diff --git a/api/core/entities/provider_configuration.py b/api/core/entities/provider_configuration.py index 95375a8cbb2..25774a7054b 100644 --- a/api/core/entities/provider_configuration.py +++ b/api/core/entities/provider_configuration.py @@ -22,7 +22,10 @@ from core.entities.provider_entities import ( SystemConfigurationStatus, ) from core.helper import encrypter -from core.helper.model_provider_cache import ProviderCredentialsCache, ProviderCredentialsCacheType +from core.helper.model_provider_cache import ( + ProviderCredentialsCache, + ProviderCredentialsCacheType, +) from core.plugin.impl.model_runtime_factory import create_model_type_instance, create_plugin_model_assembly from graphon.model_runtime.entities.model_entities import AIModelEntity, FetchFrom, ModelType from graphon.model_runtime.entities.provider_entities import ( @@ -473,6 +476,39 @@ class ProviderConfiguration(BaseModel): provider_names.append(model_provider_id.provider_name) return provider_names + def _invalidate_provider_configuration_cache( + self, + *, + provider_models: bool = False, + preferred_model_providers: bool = False, + provider_model_settings: bool = False, + provider_model_credentials: bool = False, + provider_credentials: bool = False, + provider_load_balancing_configs: bool = False, + ) -> None: + """Invalidate tenant-scoped provider snapshots after committing configuration writes.""" + from core.provider_manager import ProviderConfigurationCacheSource, ProviderManager + + sources: list[ProviderConfigurationCacheSource] = [] + if provider_models: + sources.append(ProviderConfigurationCacheSource.PROVIDER_MODELS) + if preferred_model_providers: + sources.append(ProviderConfigurationCacheSource.PREFERRED_MODEL_PROVIDERS) + if provider_model_settings: + sources.append(ProviderConfigurationCacheSource.PROVIDER_MODEL_SETTINGS) + if provider_model_credentials: + sources.append(ProviderConfigurationCacheSource.PROVIDER_MODEL_CREDENTIALS) + if provider_credentials: + sources.append(ProviderConfigurationCacheSource.PROVIDER_CREDENTIALS) + if provider_load_balancing_configs: + sources.append(ProviderConfigurationCacheSource.PROVIDER_LOAD_BALANCING_CONFIGS) + + if not sources: + logger.warning("No provider configuration cache source selected for invalidation") + return + + ProviderManager.invalidate_configurations_cache(self.tenant_id, sources=sources) + def create_provider_credential(self, credentials: dict[str, Any], credential_name: str | None): """ Add custom provider credentials. @@ -489,6 +525,7 @@ class ProviderConfiguration(BaseModel): credentials = self.validate_provider_credentials(credentials=credentials) + preferred_model_providers_changed = False with Session(db.engine) as session: provider_record = self._get_provider_record(session) try: @@ -518,7 +555,9 @@ class ProviderConfiguration(BaseModel): ) provider_model_credentials_cache.delete() - self.switch_preferred_provider_type(provider_type=ProviderType.CUSTOM, session=session) + preferred_model_providers_changed = self.switch_preferred_provider_type( + provider_type=ProviderType.CUSTOM, session=session + ) else: provider_record.is_valid = True @@ -533,12 +572,18 @@ class ProviderConfiguration(BaseModel): ) provider_model_credentials_cache.delete() - self.switch_preferred_provider_type(provider_type=ProviderType.CUSTOM, session=session) + preferred_model_providers_changed = self.switch_preferred_provider_type( + provider_type=ProviderType.CUSTOM, session=session + ) session.commit() except Exception: session.rollback() raise + self._invalidate_provider_configuration_cache( + preferred_model_providers=preferred_model_providers_changed, + provider_credentials=True, + ) def update_provider_credential( self, @@ -562,6 +607,7 @@ class ProviderConfiguration(BaseModel): credentials = self.validate_provider_credentials(credentials=credentials, credential_id=credential_id) + load_balancing_configs_changed = False with Session(db.engine) as session: provider_record = self._get_provider_record(session) stmt = select(ProviderCredential).where( @@ -588,7 +634,7 @@ class ProviderConfiguration(BaseModel): ) provider_model_credentials_cache.delete() - self._update_load_balancing_configs_with_credential( + load_balancing_configs_changed = self._update_load_balancing_configs_with_credential( credential_id=credential_id, credential_record=credential_record, credential_source=CredentialSourceType.PROVIDER, @@ -597,6 +643,10 @@ class ProviderConfiguration(BaseModel): except Exception: session.rollback() raise + self._invalidate_provider_configuration_cache( + provider_credentials=True, + provider_load_balancing_configs=load_balancing_configs_changed, + ) def _update_load_balancing_configs_with_credential( self, @@ -604,7 +654,7 @@ class ProviderConfiguration(BaseModel): credential_record: ProviderCredential | ProviderModelCredential, credential_source: str, session: Session, - ): + ) -> bool: """ Update load balancing configurations that reference the given credential_id. @@ -625,7 +675,7 @@ class ProviderConfiguration(BaseModel): load_balancing_configs = session.execute(stmt).scalars().all() if not load_balancing_configs: - return + return False # Update each load balancing config with the new credentials for lb_config in load_balancing_configs: @@ -643,6 +693,7 @@ class ProviderConfiguration(BaseModel): lb_credentials_cache.delete() session.commit() + return True def delete_provider_credential(self, credential_id: str): """ @@ -651,6 +702,8 @@ class ProviderConfiguration(BaseModel): :param credential_id: credential id :return: """ + preferred_model_providers_changed = False + load_balancing_configs_changed = False with Session(db.engine) as session: stmt = select(ProviderCredential).where( ProviderCredential.id == credential_id, @@ -671,6 +724,7 @@ class ProviderConfiguration(BaseModel): LoadBalancingModelConfig.credential_source_type == CredentialSourceType.PROVIDER, ) lb_configs_using_credential = session.execute(lb_stmt).scalars().all() + load_balancing_configs_changed = bool(lb_configs_using_credential) try: for lb_config in lb_configs_using_credential: lb_credentials_cache = ProviderCredentialsCache( @@ -703,7 +757,9 @@ class ProviderConfiguration(BaseModel): cache_type=ProviderCredentialsCacheType.PROVIDER, ) provider_model_credentials_cache.delete() - self.switch_preferred_provider_type(provider_type=ProviderType.SYSTEM, session=session) + preferred_model_providers_changed = self.switch_preferred_provider_type( + provider_type=ProviderType.SYSTEM, session=session + ) elif provider_record and provider_record.credential_id == credential_id: provider_record.credential_id = None provider_record.updated_at = naive_utc_now() @@ -714,12 +770,19 @@ class ProviderConfiguration(BaseModel): cache_type=ProviderCredentialsCacheType.PROVIDER, ) provider_model_credentials_cache.delete() - self.switch_preferred_provider_type(provider_type=ProviderType.SYSTEM, session=session) + preferred_model_providers_changed = self.switch_preferred_provider_type( + provider_type=ProviderType.SYSTEM, session=session + ) session.commit() except Exception: session.rollback() raise + self._invalidate_provider_configuration_cache( + preferred_model_providers=preferred_model_providers_changed, + provider_credentials=True, + provider_load_balancing_configs=load_balancing_configs_changed, + ) def switch_active_provider_credential(self, credential_id: str): """ @@ -728,6 +791,7 @@ class ProviderConfiguration(BaseModel): :param credential_id: credential id :return: """ + preferred_model_providers_changed = False with Session(db.engine) as session: stmt = select(ProviderCredential).where( ProviderCredential.id == credential_id, @@ -753,10 +817,14 @@ class ProviderConfiguration(BaseModel): cache_type=ProviderCredentialsCacheType.PROVIDER, ) provider_model_credentials_cache.delete() - self.switch_preferred_provider_type(ProviderType.CUSTOM, session=session) + preferred_model_providers_changed = self.switch_preferred_provider_type( + ProviderType.CUSTOM, session=session + ) except Exception: session.rollback() raise + if preferred_model_providers_changed: + self._invalidate_provider_configuration_cache(preferred_model_providers=True) def _get_custom_model_record( self, @@ -1017,6 +1085,10 @@ class ProviderConfiguration(BaseModel): except Exception: session.rollback() raise + self._invalidate_provider_configuration_cache( + provider_models=True, + provider_model_credentials=True, + ) def update_custom_model_credential( self, @@ -1053,6 +1125,7 @@ class ProviderConfiguration(BaseModel): credential_id=credential_id, ) + load_balancing_configs_changed = False with Session(db.engine) as session: provider_model_record = self._get_custom_model_record(model_type=model_type, model=model, session=session) @@ -1082,7 +1155,7 @@ class ProviderConfiguration(BaseModel): ) provider_model_credentials_cache.delete() - self._update_load_balancing_configs_with_credential( + load_balancing_configs_changed = self._update_load_balancing_configs_with_credential( credential_id=credential_id, credential_record=credential_record, credential_source=CredentialSourceType.CUSTOM_MODEL, @@ -1091,6 +1164,11 @@ class ProviderConfiguration(BaseModel): except Exception: session.rollback() raise + self._invalidate_provider_configuration_cache( + provider_models=True, + provider_model_credentials=True, + provider_load_balancing_configs=load_balancing_configs_changed, + ) def delete_custom_model_credential(self, model_type: ModelType, model: str, credential_id: str): """ @@ -1099,6 +1177,7 @@ class ProviderConfiguration(BaseModel): :param credential_id: credential id :return: """ + load_balancing_configs_changed = False with Session(db.engine) as session: stmt = select(ProviderModelCredential).where( ProviderModelCredential.id == credential_id, @@ -1118,6 +1197,7 @@ class ProviderConfiguration(BaseModel): LoadBalancingModelConfig.credential_source_type == CredentialSourceType.CUSTOM_MODEL, ) lb_configs_using_credential = session.execute(lb_stmt).scalars().all() + load_balancing_configs_changed = bool(lb_configs_using_credential) try: for lb_config in lb_configs_using_credential: @@ -1161,6 +1241,11 @@ class ProviderConfiguration(BaseModel): except Exception: session.rollback() raise + self._invalidate_provider_configuration_cache( + provider_models=True, + provider_model_credentials=True, + provider_load_balancing_configs=load_balancing_configs_changed, + ) def add_model_credential_to_model(self, model_type: ModelType, model: str, credential_id: str): """ @@ -1213,6 +1298,7 @@ class ProviderConfiguration(BaseModel): session.add(provider_model_record) session.commit() + self._invalidate_provider_configuration_cache(provider_models=True) def switch_custom_model_credential(self, model_type: ModelType, model: str, credential_id: str): """ @@ -1251,6 +1337,7 @@ class ProviderConfiguration(BaseModel): cache_type=ProviderCredentialsCacheType.MODEL, ) provider_model_credentials_cache.delete() + self._invalidate_provider_configuration_cache(provider_models=True) def delete_custom_model(self, model_type: ModelType, model: str): """ @@ -1259,6 +1346,7 @@ class ProviderConfiguration(BaseModel): :param model: model name :return: """ + provider_models_changed = False with Session(db.engine) as session: # get provider model provider_model_record = self._get_custom_model_record(model_type=model_type, model=model, session=session) @@ -1267,6 +1355,7 @@ class ProviderConfiguration(BaseModel): if provider_model_record: session.delete(provider_model_record) session.commit() + provider_models_changed = True provider_model_credentials_cache = ProviderCredentialsCache( tenant_id=self.tenant_id, @@ -1275,6 +1364,8 @@ class ProviderConfiguration(BaseModel): ) provider_model_credentials_cache.delete() + if provider_models_changed: + self._invalidate_provider_configuration_cache(provider_models=True) def _get_provider_model_setting( self, model_type: ModelType, model: str, session: Session @@ -1314,6 +1405,7 @@ class ProviderConfiguration(BaseModel): ) session.add(model_setting) session.commit() + self._invalidate_provider_configuration_cache(provider_model_settings=True) return model_setting @@ -1340,6 +1432,7 @@ class ProviderConfiguration(BaseModel): ) session.add(model_setting) session.commit() + self._invalidate_provider_configuration_cache(provider_model_settings=True) return model_setting @@ -1392,6 +1485,7 @@ class ProviderConfiguration(BaseModel): ) session.add(model_setting) session.commit() + self._invalidate_provider_configuration_cache(provider_model_settings=True) return model_setting @@ -1419,6 +1513,7 @@ class ProviderConfiguration(BaseModel): ) session.add(model_setting) session.commit() + self._invalidate_provider_configuration_cache(provider_model_settings=True) return model_setting @@ -1454,19 +1549,19 @@ class ProviderConfiguration(BaseModel): credentials=credentials or {}, ) - def switch_preferred_provider_type(self, provider_type: ProviderType, session: Session | None = None): + def switch_preferred_provider_type(self, provider_type: ProviderType, session: Session | None = None) -> bool: """ Switch preferred provider type. :param provider_type: :return: """ if provider_type == self.preferred_provider_type: - return + return False if provider_type == ProviderType.SYSTEM and not self.system_configuration.enabled: - return + return False - def _switch(s: Session): + def _switch(s: Session) -> bool: stmt = select(TenantPreferredModelProvider).where( TenantPreferredModelProvider.tenant_id == self.tenant_id, TenantPreferredModelProvider.provider_name.in_(self._get_provider_names()), @@ -1483,12 +1578,16 @@ class ProviderConfiguration(BaseModel): ) s.add(preferred_model_provider) s.commit() + return True if session: return _switch(session) else: with Session(db.engine) as session: - return _switch(session) + changed = _switch(session) + if changed: + self._invalidate_provider_configuration_cache(preferred_model_providers=True) + return changed def extract_secret_variables(self, credential_form_schemas: list[CredentialFormSchema]) -> list[str]: """ diff --git a/api/core/provider_manager.py b/api/core/provider_manager.py index 20f117c96a4..e2c710923b5 100644 --- a/api/core/provider_manager.py +++ b/api/core/provider_manager.py @@ -1,10 +1,14 @@ from __future__ import annotations import contextlib +import json +import logging from collections import defaultdict -from collections.abc import Sequence +from collections.abc import Callable, Sequence +from dataclasses import asdict, dataclass +from enum import StrEnum from json import JSONDecodeError -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING, Any, Protocol, Self from pydantic import TypeAdapter from sqlalchemy import select @@ -41,6 +45,7 @@ from graphon.model_runtime.entities.provider_entities import ( ProviderEntity, ) from graphon.model_runtime.model_providers.model_provider_factory import ModelProviderFactory +from models.enums import CredentialSourceType from models.provider import ( LoadBalancingModelConfig, Provider, @@ -59,7 +64,494 @@ if TYPE_CHECKING: from graphon.model_runtime.protocols.runtime import ModelRuntime from models.account import Account +logger = logging.getLogger(__name__) + _credentials_adapter: TypeAdapter[dict[str, Any]] = TypeAdapter(dict[str, Any]) +_PROVIDER_CONFIGURATION_CACHE_TTL_SECONDS = 300 +_PROVIDER_CONFIGURATION_CACHE_VERSION_TTL_SECONDS = 360 +_PROVIDER_CONFIGURATION_CACHE_VERSION_KEY = "provider_configurations:tenant:{tenant_id}:source:{source}:version" +_PROVIDER_CONFIGURATION_CACHE_SOURCE_KEY = "provider_configurations:tenant:{tenant_id}:source:{source}:v:{version}" + + +class ProviderConfigurationCacheSource(StrEnum): + PROVIDER_MODELS = "provider_models" + PREFERRED_MODEL_PROVIDERS = "preferred_model_providers" + PROVIDER_MODEL_SETTINGS = "provider_model_settings" + PROVIDER_MODEL_CREDENTIALS = "provider_model_credentials" + PROVIDER_CREDENTIALS = "provider_credentials" + PROVIDER_LOAD_BALANCING_CONFIGS = "provider_load_balancing_configs" + + +_PROVIDER_CONFIGURATION_SOURCES = tuple(ProviderConfigurationCacheSource) + + +class _CacheEntry(Protocol): + @classmethod + def from_cache_row(cls, row: dict[str, Any]) -> Self: ... + + def to_cache_row(self) -> dict[str, Any]: ... + + +@dataclass(frozen=True, slots=True) +class _ProviderConfigurationCacheSourceSpec[T: _CacheEntry]: + name: ProviderConfigurationCacheSource + entry_cls: type[T] + load_records: Callable[[str], list[T]] + + +@dataclass(frozen=True, slots=True) +class _ProviderModelCacheEntry: + id: str + provider_name: str + model_name: str + model_type: ModelType + credential_id: str | None + credential_name: str | None + encrypted_config: str | None + + @classmethod + def from_record(cls, record: ProviderModel) -> _ProviderModelCacheEntry: + credential = record.__dict__.get("credential") + return cls( + id=record.id, + provider_name=record.provider_name, + model_name=record.model_name, + model_type=record.model_type, + credential_id=record.credential_id, + credential_name=credential.credential_name if credential else None, + encrypted_config=credential.encrypted_config if credential else None, + ) + + @classmethod + def from_cache_row(cls, row: dict[str, Any]) -> _ProviderModelCacheEntry: + return cls( + id=row["id"], + provider_name=row["provider_name"], + model_name=row["model_name"], + model_type=ModelType(row["model_type"]), + credential_id=row.get("credential_id"), + credential_name=row.get("credential_name"), + encrypted_config=row.get("encrypted_config"), + ) + + def to_cache_row(self) -> dict[str, Any]: + row = asdict(self) + row["model_type"] = self.model_type.value + return row + + +@dataclass(frozen=True, slots=True) +class _TenantPreferredModelProviderCacheEntry: + provider_name: str + preferred_provider_type: ProviderType + + @classmethod + def from_record(cls, record: TenantPreferredModelProvider) -> _TenantPreferredModelProviderCacheEntry: + return cls( + provider_name=record.provider_name, + preferred_provider_type=record.preferred_provider_type, + ) + + @classmethod + def from_cache_row(cls, row: dict[str, Any]) -> _TenantPreferredModelProviderCacheEntry: + return cls( + provider_name=row["provider_name"], + preferred_provider_type=ProviderType(row["preferred_provider_type"]), + ) + + def to_cache_row(self) -> dict[str, Any]: + return { + "provider_name": self.provider_name, + "preferred_provider_type": self.preferred_provider_type.value, + } + + +@dataclass(frozen=True, slots=True) +class _ProviderModelSettingCacheEntry: + provider_name: str + model_name: str + model_type: ModelType + enabled: bool + load_balancing_enabled: bool + + @classmethod + def from_record(cls, record: ProviderModelSetting) -> _ProviderModelSettingCacheEntry: + return cls( + provider_name=record.provider_name, + model_name=record.model_name, + model_type=record.model_type, + enabled=record.enabled, + load_balancing_enabled=record.load_balancing_enabled, + ) + + @classmethod + def from_cache_row(cls, row: dict[str, Any]) -> _ProviderModelSettingCacheEntry: + return cls( + provider_name=row["provider_name"], + model_name=row["model_name"], + model_type=ModelType(row["model_type"]), + enabled=row["enabled"], + load_balancing_enabled=row["load_balancing_enabled"], + ) + + def to_cache_row(self) -> dict[str, Any]: + return { + "provider_name": self.provider_name, + "model_name": self.model_name, + "model_type": self.model_type.value, + "enabled": self.enabled, + "load_balancing_enabled": self.load_balancing_enabled, + } + + +@dataclass(frozen=True, slots=True) +class _ProviderModelCredentialCacheEntry: + id: str + provider_name: str + model_name: str + model_type: ModelType + credential_name: str + + @classmethod + def from_record(cls, record: ProviderModelCredential) -> _ProviderModelCredentialCacheEntry: + return cls( + id=record.id, + provider_name=record.provider_name, + model_name=record.model_name, + model_type=record.model_type, + credential_name=record.credential_name, + ) + + @classmethod + def from_cache_row(cls, row: dict[str, Any]) -> _ProviderModelCredentialCacheEntry: + return cls( + id=row["id"], + provider_name=row["provider_name"], + model_name=row["model_name"], + model_type=ModelType(row["model_type"]), + credential_name=row["credential_name"], + ) + + def to_cache_row(self) -> dict[str, Any]: + return { + "id": self.id, + "provider_name": self.provider_name, + "model_name": self.model_name, + "model_type": self.model_type.value, + "credential_name": self.credential_name, + } + + +@dataclass(frozen=True, slots=True) +class _ProviderCredentialCacheEntry: + id: str + provider_name: str + credential_name: str + + @classmethod + def from_record(cls, record: ProviderCredential) -> _ProviderCredentialCacheEntry: + return cls( + id=record.id, + provider_name=record.provider_name, + credential_name=record.credential_name, + ) + + @classmethod + def from_cache_row(cls, row: dict[str, Any]) -> _ProviderCredentialCacheEntry: + return cls( + id=row["id"], + provider_name=row["provider_name"], + credential_name=row["credential_name"], + ) + + def to_cache_row(self) -> dict[str, Any]: + return asdict(self) + + +@dataclass(frozen=True, slots=True) +class _LoadBalancingModelConfigCacheEntry: + id: str + tenant_id: str + provider_name: str + model_name: str + model_type: ModelType + name: str + encrypted_config: str | None + credential_id: str | None + credential_source_type: CredentialSourceType | None + enabled: bool + + @classmethod + def from_record(cls, record: LoadBalancingModelConfig) -> _LoadBalancingModelConfigCacheEntry: + return cls( + id=record.id, + tenant_id=record.tenant_id, + provider_name=record.provider_name, + model_name=record.model_name, + model_type=record.model_type, + name=record.name, + encrypted_config=record.encrypted_config, + credential_id=record.credential_id, + credential_source_type=record.credential_source_type, + enabled=record.enabled, + ) + + @classmethod + def from_cache_row(cls, row: dict[str, Any]) -> _LoadBalancingModelConfigCacheEntry: + return cls( + id=row["id"], + tenant_id=row["tenant_id"], + provider_name=row["provider_name"], + model_name=row["model_name"], + model_type=ModelType(row["model_type"]), + name=row["name"], + encrypted_config=row.get("encrypted_config"), + credential_id=row.get("credential_id"), + credential_source_type=CredentialSourceType(row["credential_source_type"]) + if row.get("credential_source_type") + else None, + enabled=row["enabled"], + ) + + def to_cache_row(self) -> dict[str, Any]: + return { + "id": self.id, + "tenant_id": self.tenant_id, + "provider_name": self.provider_name, + "model_name": self.model_name, + "model_type": self.model_type.value, + "name": self.name, + "encrypted_config": self.encrypted_config, + "credential_id": self.credential_id, + "credential_source_type": self.credential_source_type.value if self.credential_source_type else None, + "enabled": self.enabled, + } + + +class _ProviderConfigurationSourceCache: + """Redis-backed cache for tenant provider DB cache entries. + + The assembled ``ProviderConfigurations`` object is intentionally not cached + here because it carries request-scoped runtime bindings. Cache only the DB + rows that are stable enough to reuse across processes, then let each + ``ProviderManager`` assemble and bind fresh runtime-aware entities. + """ + + @classmethod + def get_records[T: _CacheEntry]( + cls, + *, + tenant_id: str, + source: ProviderConfigurationCacheSource, + entry_cls: type[T], + ) -> tuple[list[T] | None, str | None]: + version: str | None = None + try: + version = cls._get_version(tenant_id=tenant_id, source=source) + cache_key = cls._source_key(tenant_id=tenant_id, source=source, version=version) + cached_records = redis_client.get(cache_key) + if cached_records is None: + return None, version + + cached_text = cached_records.decode("utf-8") if isinstance(cached_records, bytes) else cached_records + rows = json.loads(cached_text) + if not isinstance(rows, list): + return None, version + + return [entry_cls.from_cache_row(row) for row in rows if isinstance(row, dict)], version + except Exception: + logger.warning("Failed to read provider configuration source cache", exc_info=True) + return None, version + + @classmethod + def set_records( + cls, + *, + tenant_id: str, + source: ProviderConfigurationCacheSource, + records: Sequence[_CacheEntry], + expected_version: str | None = None, + ) -> None: + try: + version = cls._get_version(tenant_id=tenant_id, source=source) + if expected_version is not None and version != expected_version: + return + cache_key = cls._source_key(tenant_id=tenant_id, source=source, version=version) + rows = [record.to_cache_row() for record in records] + redis_client.setex(cache_key, _PROVIDER_CONFIGURATION_CACHE_TTL_SECONDS, json.dumps(rows)) + except Exception: + logger.warning("Failed to write provider configuration source cache", exc_info=True) + + @classmethod + def invalidate_tenant( + cls, + tenant_id: str, + sources: Sequence[ProviderConfigurationCacheSource] | None = None, + ) -> None: + try: + if sources is None: + sources = _PROVIDER_CONFIGURATION_SOURCES + for source in sources: + version_key = _PROVIDER_CONFIGURATION_CACHE_VERSION_KEY.format(tenant_id=tenant_id, source=source.value) + redis_client.incr(version_key) + redis_client.expire(version_key, _PROVIDER_CONFIGURATION_CACHE_VERSION_TTL_SECONDS) + except Exception: + logger.warning("Failed to invalidate provider configuration source cache", exc_info=True) + + @classmethod + def _get_version(cls, *, tenant_id: str, source: ProviderConfigurationCacheSource) -> str: + version_key = _PROVIDER_CONFIGURATION_CACHE_VERSION_KEY.format(tenant_id=tenant_id, source=source.value) + version = redis_client.get(version_key) + if version is None: + redis_client.set(version_key, "0", ex=_PROVIDER_CONFIGURATION_CACHE_VERSION_TTL_SECONDS) + return "0" + redis_client.expire(version_key, _PROVIDER_CONFIGURATION_CACHE_VERSION_TTL_SECONDS) + return version.decode("utf-8") if isinstance(version, bytes) else str(version) + + @staticmethod + def _source_key(*, tenant_id: str, source: ProviderConfigurationCacheSource, version: str) -> str: + return _PROVIDER_CONFIGURATION_CACHE_SOURCE_KEY.format( + tenant_id=tenant_id, + source=source.value, + version=version, + ) + + +def _get_cached_or_load_records[T: _CacheEntry]( + *, + tenant_id: str, + cache_source: _ProviderConfigurationCacheSourceSpec[T], +) -> list[T]: + cached_records, cache_version = _ProviderConfigurationSourceCache.get_records( + tenant_id=tenant_id, + source=cache_source.name, + entry_cls=cache_source.entry_cls, + ) + if cached_records is not None: + return cached_records + + records = cache_source.load_records(tenant_id) + _ProviderConfigurationSourceCache.set_records( + tenant_id=tenant_id, + source=cache_source.name, + records=records, + expected_version=cache_version, + ) + return records + + +def _attach_active_credentials( + *, + session: Any, + records: Sequence[Provider | ProviderModel], + credential_model_cls: type[ProviderCredential | ProviderModelCredential], +) -> None: + credential_ids = [record.credential_id for record in records if getattr(record, "credential_id", None)] + if not credential_ids: + return + + credentials = session.scalars(select(credential_model_cls).where(credential_model_cls.id.in_(credential_ids))).all() + credential_by_id = {credential.id: credential for credential in credentials} + for record in records: + if getattr(record, "credential_id", None): + record.__dict__["credential"] = credential_by_id.get(record.credential_id) + + +def _load_provider_model_cache_entries(tenant_id: str) -> list[_ProviderModelCacheEntry]: + with session_factory.create_session() as session: + stmt = select(ProviderModel).where(ProviderModel.tenant_id == tenant_id, ProviderModel.is_valid == True) + provider_models = list(session.scalars(stmt)) + _attach_active_credentials( + session=session, + records=provider_models, + credential_model_cls=ProviderModelCredential, + ) + return [_ProviderModelCacheEntry.from_record(provider_model) for provider_model in provider_models] + + +def _load_preferred_model_provider_cache_entries(tenant_id: str) -> list[_TenantPreferredModelProviderCacheEntry]: + with session_factory.create_session() as session: + stmt = select(TenantPreferredModelProvider).where(TenantPreferredModelProvider.tenant_id == tenant_id) + return [ + _TenantPreferredModelProviderCacheEntry.from_record(preferred_model_provider) + for preferred_model_provider in session.scalars(stmt) + ] + + +def _load_provider_model_setting_cache_entries(tenant_id: str) -> list[_ProviderModelSettingCacheEntry]: + with session_factory.create_session() as session: + stmt = select(ProviderModelSetting).where(ProviderModelSetting.tenant_id == tenant_id) + return [ + _ProviderModelSettingCacheEntry.from_record(provider_model_setting) + for provider_model_setting in session.scalars(stmt) + ] + + +def _load_provider_model_credential_cache_entries(tenant_id: str) -> list[_ProviderModelCredentialCacheEntry]: + with session_factory.create_session() as session: + stmt = ( + select(ProviderModelCredential) + .where(ProviderModelCredential.tenant_id == tenant_id) + .order_by(ProviderModelCredential.created_at.desc()) + ) + return [ + _ProviderModelCredentialCacheEntry.from_record(provider_model_credential) + for provider_model_credential in session.scalars(stmt) + ] + + +def _load_provider_credential_cache_entries(tenant_id: str) -> list[_ProviderCredentialCacheEntry]: + with session_factory.create_session() as session: + stmt = ( + select(ProviderCredential) + .where(ProviderCredential.tenant_id == tenant_id) + .order_by(ProviderCredential.created_at.desc()) + ) + return [ + _ProviderCredentialCacheEntry.from_record(provider_credential) + for provider_credential in session.scalars(stmt) + ] + + +def _load_provider_load_balancing_config_cache_entries(tenant_id: str) -> list[_LoadBalancingModelConfigCacheEntry]: + with session_factory.create_session() as session: + stmt = select(LoadBalancingModelConfig).where(LoadBalancingModelConfig.tenant_id == tenant_id) + return [ + _LoadBalancingModelConfigCacheEntry.from_record(load_balancing_model_config) + for load_balancing_model_config in session.scalars(stmt) + ] + + +_PROVIDER_MODELS_CACHE_SOURCE = _ProviderConfigurationCacheSourceSpec( + name=ProviderConfigurationCacheSource.PROVIDER_MODELS, + entry_cls=_ProviderModelCacheEntry, + load_records=_load_provider_model_cache_entries, +) +_PREFERRED_MODEL_PROVIDERS_CACHE_SOURCE = _ProviderConfigurationCacheSourceSpec( + name=ProviderConfigurationCacheSource.PREFERRED_MODEL_PROVIDERS, + entry_cls=_TenantPreferredModelProviderCacheEntry, + load_records=_load_preferred_model_provider_cache_entries, +) +_PROVIDER_MODEL_SETTINGS_CACHE_SOURCE = _ProviderConfigurationCacheSourceSpec( + name=ProviderConfigurationCacheSource.PROVIDER_MODEL_SETTINGS, + entry_cls=_ProviderModelSettingCacheEntry, + load_records=_load_provider_model_setting_cache_entries, +) +_PROVIDER_MODEL_CREDENTIALS_CACHE_SOURCE = _ProviderConfigurationCacheSourceSpec( + name=ProviderConfigurationCacheSource.PROVIDER_MODEL_CREDENTIALS, + entry_cls=_ProviderModelCredentialCacheEntry, + load_records=_load_provider_model_credential_cache_entries, +) +_PROVIDER_CREDENTIALS_CACHE_SOURCE = _ProviderConfigurationCacheSourceSpec( + name=ProviderConfigurationCacheSource.PROVIDER_CREDENTIALS, + entry_cls=_ProviderCredentialCacheEntry, + load_records=_load_provider_credential_cache_entries, +) +_PROVIDER_LOAD_BALANCING_CONFIGS_CACHE_SOURCE = _ProviderConfigurationCacheSourceSpec( + name=ProviderConfigurationCacheSource.PROVIDER_LOAD_BALANCING_CONFIGS, + entry_cls=_LoadBalancingModelConfigCacheEntry, + load_records=_load_provider_load_balancing_config_cache_entries, +) class ProviderManager: @@ -98,6 +590,14 @@ class ProviderManager: self._configurations_cache.pop(tenant_id, None) + @staticmethod + def invalidate_configurations_cache( + tenant_id: str, + sources: Sequence[ProviderConfigurationCacheSource] | None = None, + ) -> None: + """Invalidate cross-process provider configuration source cache for a tenant.""" + _ProviderConfigurationSourceCache.invalidate_tenant(tenant_id, sources=sources) + def get_configurations(self, tenant_id: str) -> ProviderConfigurations: """ Get model provider configurations. @@ -192,6 +692,9 @@ class ProviderManager: # Get All provider model credentials provider_name_to_provider_model_credentials_dict = self._get_all_provider_model_credentials(tenant_id) + # Get All provider credentials + provider_name_to_provider_credentials_dict = self._get_all_provider_credentials(tenant_id) + provider_configurations = ProviderConfigurations(tenant_id=tenant_id) # Construct ProviderConfiguration objects for each provider @@ -224,7 +727,12 @@ class ProviderManager: # Convert to custom configuration custom_configuration = self._to_custom_configuration( - tenant_id, provider_entity, provider_records, provider_model_records, provider_model_credentials + tenant_id, + provider_entity, + provider_records, + provider_model_records, + provider_model_credentials, + provider_name_to_provider_credentials_dict, ) # Convert to system configuration @@ -448,84 +956,115 @@ class ProviderManager: provider_name_to_provider_records_dict = defaultdict(list) with session_factory.create_session() as session: stmt = select(Provider).where(Provider.tenant_id == tenant_id, Provider.is_valid == True) - providers = session.scalars(stmt) + providers = list(session.scalars(stmt)) + _attach_active_credentials( + session=session, + records=providers, + credential_model_cls=ProviderCredential, + ) for provider in providers: # Use provider name with prefix after the data migration provider_name_to_provider_records_dict[str(ModelProviderID(provider.provider_name))].append(provider) return provider_name_to_provider_records_dict @staticmethod - def _get_all_provider_models(tenant_id: str) -> dict[str, list[ProviderModel]]: + def _get_all_provider_models(tenant_id: str) -> dict[str, list[_ProviderModelCacheEntry]]: """ Get all provider model records of the workspace. :param tenant_id: workspace id :return: """ + provider_models = _get_cached_or_load_records( + tenant_id=tenant_id, + cache_source=_PROVIDER_MODELS_CACHE_SOURCE, + ) + provider_name_to_provider_model_records_dict = defaultdict(list) - with session_factory.create_session() as session: - stmt = select(ProviderModel).where(ProviderModel.tenant_id == tenant_id, ProviderModel.is_valid == True) - provider_models = session.scalars(stmt) - for provider_model in provider_models: - provider_name_to_provider_model_records_dict[provider_model.provider_name].append(provider_model) + for provider_model in provider_models: + provider_name_to_provider_model_records_dict[provider_model.provider_name].append(provider_model) return provider_name_to_provider_model_records_dict @staticmethod - def _get_all_preferred_model_providers(tenant_id: str) -> dict[str, TenantPreferredModelProvider]: + def _get_all_preferred_model_providers(tenant_id: str) -> dict[str, _TenantPreferredModelProviderCacheEntry]: """ Get All preferred provider types of the workspace. :param tenant_id: workspace id :return: """ - provider_name_to_preferred_provider_type_records_dict = {} - with session_factory.create_session() as session: - stmt = select(TenantPreferredModelProvider).where(TenantPreferredModelProvider.tenant_id == tenant_id) - preferred_provider_types = session.scalars(stmt) - provider_name_to_preferred_provider_type_records_dict = { - preferred_provider_type.provider_name: preferred_provider_type - for preferred_provider_type in preferred_provider_types - } - return provider_name_to_preferred_provider_type_records_dict + preferred_provider_types = _get_cached_or_load_records( + tenant_id=tenant_id, + cache_source=_PREFERRED_MODEL_PROVIDERS_CACHE_SOURCE, + ) + + return { + preferred_provider_type.provider_name: preferred_provider_type + for preferred_provider_type in preferred_provider_types + } @staticmethod - def _get_all_provider_model_settings(tenant_id: str) -> dict[str, list[ProviderModelSetting]]: + def _get_all_provider_model_settings(tenant_id: str) -> dict[str, list[_ProviderModelSettingCacheEntry]]: """ Get All provider model settings of the workspace. :param tenant_id: workspace id :return: """ + provider_model_settings = _get_cached_or_load_records( + tenant_id=tenant_id, + cache_source=_PROVIDER_MODEL_SETTINGS_CACHE_SOURCE, + ) + provider_name_to_provider_model_settings_dict = defaultdict(list) - with session_factory.create_session() as session: - stmt = select(ProviderModelSetting).where(ProviderModelSetting.tenant_id == tenant_id) - provider_model_settings = session.scalars(stmt) - for provider_model_setting in provider_model_settings: - provider_name_to_provider_model_settings_dict[provider_model_setting.provider_name].append( - provider_model_setting - ) + for provider_model_setting in provider_model_settings: + provider_name_to_provider_model_settings_dict[provider_model_setting.provider_name].append( + provider_model_setting + ) return provider_name_to_provider_model_settings_dict @staticmethod - def _get_all_provider_model_credentials(tenant_id: str) -> dict[str, list[ProviderModelCredential]]: + def _get_all_provider_model_credentials(tenant_id: str) -> dict[str, list[_ProviderModelCredentialCacheEntry]]: """ Get All provider model credentials of the workspace. :param tenant_id: workspace id :return: """ + provider_model_credentials = _get_cached_or_load_records( + tenant_id=tenant_id, + cache_source=_PROVIDER_MODEL_CREDENTIALS_CACHE_SOURCE, + ) + provider_name_to_provider_model_credentials_dict = defaultdict(list) - with session_factory.create_session() as session: - stmt = select(ProviderModelCredential).where(ProviderModelCredential.tenant_id == tenant_id) - provider_model_credentials = session.scalars(stmt) - for provider_model_credential in provider_model_credentials: - provider_name_to_provider_model_credentials_dict[provider_model_credential.provider_name].append( - provider_model_credential - ) + for provider_model_credential in provider_model_credentials: + provider_name_to_provider_model_credentials_dict[provider_model_credential.provider_name].append( + provider_model_credential + ) return provider_name_to_provider_model_credentials_dict @staticmethod - def _get_all_provider_load_balancing_configs(tenant_id: str) -> dict[str, list[LoadBalancingModelConfig]]: + def _get_all_provider_credentials(tenant_id: str) -> dict[str, list[_ProviderCredentialCacheEntry]]: + """ + Get All provider credentials of the workspace. + + :param tenant_id: workspace id + :return: + """ + provider_credentials = _get_cached_or_load_records( + tenant_id=tenant_id, + cache_source=_PROVIDER_CREDENTIALS_CACHE_SOURCE, + ) + + provider_name_to_provider_credentials_dict = defaultdict(list) + for provider_credential in provider_credentials: + provider_name_to_provider_credentials_dict[provider_credential.provider_name].append(provider_credential) + return provider_name_to_provider_credentials_dict + + @staticmethod + def _get_all_provider_load_balancing_configs( + tenant_id: str, + ) -> dict[str, list[_LoadBalancingModelConfigCacheEntry]]: """ Get All provider load balancing configs of the workspace. @@ -546,14 +1085,16 @@ class ProviderManager: if not model_load_balancing_enabled: return {} + provider_load_balancing_configs = _get_cached_or_load_records( + tenant_id=tenant_id, + cache_source=_PROVIDER_LOAD_BALANCING_CONFIGS_CACHE_SOURCE, + ) + provider_name_to_provider_load_balancing_model_configs_dict = defaultdict(list) - with session_factory.create_session() as session: - stmt = select(LoadBalancingModelConfig).where(LoadBalancingModelConfig.tenant_id == tenant_id) - provider_load_balancing_configs = session.scalars(stmt) - for provider_load_balancing_config in provider_load_balancing_configs: - provider_name_to_provider_load_balancing_model_configs_dict[ - provider_load_balancing_config.provider_name - ].append(provider_load_balancing_config) + for provider_load_balancing_config in provider_load_balancing_configs: + provider_name_to_provider_load_balancing_model_configs_dict[ + provider_load_balancing_config.provider_name + ].append(provider_load_balancing_config) return provider_name_to_provider_load_balancing_model_configs_dict @@ -722,8 +1263,9 @@ class ProviderManager: tenant_id: str, provider_entity: ProviderEntity, provider_records: list[Provider], - provider_model_records: list[ProviderModel], - provider_model_credentials: list[ProviderModelCredential], + provider_model_records: list[_ProviderModelCacheEntry], + provider_model_credentials: list[_ProviderModelCredentialCacheEntry], + provider_credentials_by_name: dict[str, list[_ProviderCredentialCacheEntry]], ) -> CustomConfiguration: """ Convert to custom configuration. @@ -736,7 +1278,10 @@ class ProviderManager: """ # Get custom provider configuration custom_provider_configuration = self._get_custom_provider_configuration( - tenant_id, provider_entity, provider_records + tenant_id, + provider_entity, + provider_records, + provider_credentials_by_name, ) # Get custom models which have not been added to the model list yet @@ -758,7 +1303,11 @@ class ProviderManager: ) def _get_custom_provider_configuration( - self, tenant_id: str, provider_entity: ProviderEntity, provider_records: list[Provider] + self, + tenant_id: str, + provider_entity: ProviderEntity, + provider_records: list[Provider], + provider_credentials_by_name: dict[str, list[_ProviderCredentialCacheEntry]], ) -> CustomProviderConfiguration | None: """Get custom provider configuration.""" # Find custom provider record (non-system) @@ -790,13 +1339,29 @@ class ProviderManager: credentials=provider_credentials, current_credential_name=custom_provider_record.credential_name, current_credential_id=custom_provider_record.credential_id, - available_credentials=self.get_provider_available_credentials( - tenant_id, custom_provider_record.provider_name + available_credentials=self._get_provider_available_credentials_from_records( + custom_provider_record.provider_name, + provider_credentials_by_name, ), ) + @staticmethod + def _get_provider_available_credentials_from_records( + provider_name: str, + provider_credentials_by_name: dict[str, list[_ProviderCredentialCacheEntry]], + ) -> list[CredentialConfiguration]: + available_credentials: list[CredentialConfiguration] = [] + for candidate_provider_name in ProviderManager._get_provider_names(provider_name): + available_credentials.extend( + CredentialConfiguration(credential_id=credential.id, credential_name=credential.credential_name) + for credential in provider_credentials_by_name.get(candidate_provider_name, []) + ) + return available_credentials + def _get_can_added_models( - self, provider_model_records: list[ProviderModel], all_model_credentials: Sequence[ProviderModelCredential] + self, + provider_model_records: list[_ProviderModelCacheEntry], + all_model_credentials: Sequence[_ProviderModelCredentialCacheEntry], ) -> list[dict]: """Get the custom models and credentials from enterprise version which haven't add to the model list""" existing_model_set = {(record.model_name, record.model_type) for record in provider_model_records} @@ -829,9 +1394,9 @@ class ProviderManager: self, tenant_id: str, provider_entity: ProviderEntity, - provider_model_records: list[ProviderModel], + provider_model_records: list[_ProviderModelCacheEntry], can_added_models: list[dict], - all_model_credentials: Sequence[ProviderModelCredential], + all_model_credentials: Sequence[_ProviderModelCredentialCacheEntry], ) -> list[CustomModelConfiguration]: """Get custom model configurations.""" # Get model credential secret variables @@ -1151,8 +1716,8 @@ class ProviderManager: def _to_model_settings( self, provider_entity: ProviderEntity, - provider_model_settings: list[ProviderModelSetting] | None = None, - load_balancing_model_configs: list[LoadBalancingModelConfig] | None = None, + provider_model_settings: list[_ProviderModelSettingCacheEntry] | None = None, + load_balancing_model_configs: list[_LoadBalancingModelConfigCacheEntry] | None = None, ) -> list[ModelSettings]: """ Convert to model settings. diff --git a/api/services/model_load_balancing_service.py b/api/services/model_load_balancing_service.py index 46bf24fffbf..2a9094a35f2 100644 --- a/api/services/model_load_balancing_service.py +++ b/api/services/model_load_balancing_service.py @@ -7,10 +7,13 @@ from sqlalchemy import or_, select from constants import HIDDEN_VALUE from core.entities.provider_configuration import ProviderConfiguration from core.helper import encrypter -from core.helper.model_provider_cache import ProviderCredentialsCache, ProviderCredentialsCacheType +from core.helper.model_provider_cache import ( + ProviderCredentialsCache, + ProviderCredentialsCacheType, +) from core.model_manager import LBModelManager from core.plugin.impl.model_runtime_factory import create_plugin_model_assembly, create_plugin_provider_manager -from core.provider_manager import ProviderManager +from core.provider_manager import ProviderConfigurationCacheSource, ProviderManager from extensions.ext_database import db from graphon.model_runtime.entities.model_entities import ModelType from graphon.model_runtime.entities.provider_entities import ( @@ -313,6 +316,10 @@ class ModelLoadBalancingService: ) db.session.add(inherit_config) db.session.commit() + ProviderManager.invalidate_configurations_cache( + tenant_id, + sources=(ProviderConfigurationCacheSource.PROVIDER_LOAD_BALANCING_CONFIGS,), + ) return inherit_config @@ -434,6 +441,10 @@ class ModelLoadBalancingService: load_balancing_config.enabled = enabled load_balancing_config.updated_at = naive_utc_now() db.session.commit() + ProviderManager.invalidate_configurations_cache( + tenant_id, + sources=(ProviderConfigurationCacheSource.PROVIDER_LOAD_BALANCING_CONFIGS,), + ) self._clear_credentials_cache(tenant_id, config_id) else: @@ -487,12 +498,20 @@ class ModelLoadBalancingService: db.session.add(load_balancing_model_config) db.session.commit() + ProviderManager.invalidate_configurations_cache( + tenant_id, + sources=(ProviderConfigurationCacheSource.PROVIDER_LOAD_BALANCING_CONFIGS,), + ) # get deleted config ids deleted_config_ids = set(current_load_balancing_configs_dict.keys()) - updated_config_ids for config_id in deleted_config_ids: db.session.delete(current_load_balancing_configs_dict[config_id]) db.session.commit() + ProviderManager.invalidate_configurations_cache( + tenant_id, + sources=(ProviderConfigurationCacheSource.PROVIDER_LOAD_BALANCING_CONFIGS,), + ) self._clear_credentials_cache(tenant_id, config_id) diff --git a/api/tests/unit_tests/core/entities/test_entities_provider_configuration.py b/api/tests/unit_tests/core/entities/test_entities_provider_configuration.py index bb473739c87..23d8fd28df0 100644 --- a/api/tests/unit_tests/core/entities/test_entities_provider_configuration.py +++ b/api/tests/unit_tests/core/entities/test_entities_provider_configuration.py @@ -394,13 +394,15 @@ def test_switch_preferred_provider_type_returns_early_when_no_change_or_unsuppor configuration = _build_provider_configuration() with patch("core.entities.provider_configuration.Session") as mock_session_cls: - configuration.switch_preferred_provider_type(ProviderType.SYSTEM) + changed = configuration.switch_preferred_provider_type(ProviderType.SYSTEM) + assert changed is False mock_session_cls.assert_not_called() configuration.preferred_provider_type = ProviderType.CUSTOM configuration.system_configuration.enabled = False with patch("core.entities.provider_configuration.Session") as mock_session_cls: - configuration.switch_preferred_provider_type(ProviderType.SYSTEM) + changed = configuration.switch_preferred_provider_type(ProviderType.SYSTEM) + assert changed is False mock_session_cls.assert_not_called() @@ -411,10 +413,13 @@ def test_switch_preferred_provider_type_updates_existing_record_with_session() - existing_record = SimpleNamespace(preferred_provider_type="custom") session.execute.return_value.scalars.return_value.first.return_value = existing_record - configuration.switch_preferred_provider_type(ProviderType.SYSTEM, session=session) + with patch.object(ProviderConfiguration, "_invalidate_provider_configuration_cache") as mock_invalidate: + changed = configuration.switch_preferred_provider_type(ProviderType.SYSTEM, session=session) + assert changed is True assert existing_record.preferred_provider_type == ProviderType.SYSTEM session.commit.assert_called_once() + mock_invalidate.assert_not_called() def test_switch_preferred_provider_type_creates_record_when_missing() -> None: @@ -423,10 +428,13 @@ def test_switch_preferred_provider_type_creates_record_when_missing() -> None: session = Mock() session.execute.return_value.scalars.return_value.first.return_value = None - configuration.switch_preferred_provider_type(ProviderType.CUSTOM, session=session) + with patch.object(ProviderConfiguration, "_invalidate_provider_configuration_cache") as mock_invalidate: + changed = configuration.switch_preferred_provider_type(ProviderType.CUSTOM, session=session) + assert changed is True assert session.add.call_count == 1 session.commit.assert_called_once() + mock_invalidate.assert_not_called() def test_get_model_type_instance_and_schema_delegate_to_factory() -> None: @@ -1022,13 +1030,14 @@ def test_update_load_balancing_configs_updates_all_matching_configs() -> None: credential_record = SimpleNamespace(encrypted_config='{"api_key":"enc"}', credential_name="API KEY 3") with patch("core.entities.provider_configuration.ProviderCredentialsCache") as mock_cache: - configuration._update_load_balancing_configs_with_credential( + changed = configuration._update_load_balancing_configs_with_credential( credential_id="cred-1", credential_record=credential_record, credential_source=CredentialSourceType.PROVIDER, session=session, ) + assert changed is True assert lb_config.encrypted_config == '{"api_key":"enc"}' assert lb_config.name == "API KEY 3" mock_cache.return_value.delete.assert_called_once() @@ -1040,13 +1049,14 @@ def test_update_load_balancing_configs_returns_when_no_matching_configs() -> Non session = Mock() session.execute.return_value.scalars.return_value.all.return_value = [] - configuration._update_load_balancing_configs_with_credential( + changed = configuration._update_load_balancing_configs_with_credential( credential_id="cred-1", credential_record=SimpleNamespace(encrypted_config="{}", credential_name="Main"), credential_source=CredentialSourceType.PROVIDER, session=session, ) + assert changed is False session.commit.assert_not_called() @@ -1478,12 +1488,15 @@ def test_model_load_balancing_enable_disable_and_switch_preferred_provider_type_ switch_session = Mock() with _patched_session(switch_session): switch_session.execute.return_value.scalars.return_value.first.return_value = None - configuration.switch_preferred_provider_type(ProviderType.CUSTOM) + with patch.object(ProviderConfiguration, "_invalidate_provider_configuration_cache") as mock_invalidate: + changed = configuration.switch_preferred_provider_type(ProviderType.CUSTOM) + assert changed is True assert any( call.args and call.args[0].__class__.__name__ == "TenantPreferredModelProvider" for call in switch_session.add.call_args_list ) switch_session.commit.assert_called() + mock_invalidate.assert_called_once_with(preferred_model_providers=True) def test_system_and_custom_provider_model_helpers_cover_remaining_skip_paths() -> None: diff --git a/api/tests/unit_tests/core/test_provider_manager.py b/api/tests/unit_tests/core/test_provider_manager.py index e84fcba3d94..128eebdd5af 100644 --- a/api/tests/unit_tests/core/test_provider_manager.py +++ b/api/tests/unit_tests/core/test_provider_manager.py @@ -5,10 +5,17 @@ import pytest from pytest_mock import MockerFixture from core.entities.provider_entities import ModelSettings -from core.provider_manager import ProviderManager +from core.provider_manager import ProviderConfigurationCacheSource, ProviderManager from graphon.model_runtime.entities.common_entities import I18nObject from graphon.model_runtime.entities.model_entities import ModelType -from models.provider import LoadBalancingModelConfig, ProviderModelSetting, TenantDefaultModel +from models.provider import ( + LoadBalancingModelConfig, + Provider, + ProviderCredential, + ProviderModelSetting, + ProviderType, + TenantDefaultModel, +) from models.provider_ids import ModelProviderID @@ -23,6 +30,43 @@ def _build_session_context(session: Mock) -> MagicMock: return session_cm +class _FakeRedis: + def __init__(self) -> None: + self.store: dict[str, str] = {} + self.expirations: dict[str, int] = {} + + def get(self, key: str): + return self.store.get(key) + + def set(self, key: str, value: str, *, ex: int | None = None) -> None: + self.store[key] = value + if ex is not None: + self.expirations[key] = ex + + def setex(self, key: str, time: int, value: str) -> None: + self.store[key] = value + self.expirations[key] = time + + def incr(self, key: str) -> int: + value = int(self.store.get(key, "0")) + 1 + self.store[key] = str(value) + return value + + def expire(self, key: str, time: int) -> None: + self.expirations[key] = time + + +class _FakeScalarResult: + def __init__(self, values: list[object]) -> None: + self._values = values + + def __iter__(self): + return iter(self._values) + + def all(self) -> list[object]: + return self._values + + @pytest.fixture def mock_provider_entity(): mock_entity = Mock() @@ -309,6 +353,7 @@ def test_get_configurations_uses_injected_runtime_and_adds_provider_aliases(mock patch.object(manager, "_get_all_provider_model_settings", return_value={}), patch.object(manager, "_get_all_provider_load_balancing_configs", return_value={}), patch.object(manager, "_get_all_provider_model_credentials", return_value={}), + patch.object(manager, "_get_all_provider_credentials", return_value={}), patch("core.provider_manager.ModelProviderFactory") as mock_factory_cls, ): mock_factory_cls.return_value.get_providers.return_value = [] @@ -361,6 +406,7 @@ def test_get_configurations_binds_manager_runtime_to_provider_configuration( patch.object(manager, "_get_all_provider_model_settings", return_value={}), patch.object(manager, "_get_all_provider_load_balancing_configs", return_value={}), patch.object(manager, "_get_all_provider_model_credentials", return_value={}), + patch.object(manager, "_get_all_provider_credentials", return_value={}), patch.object(manager, "_to_custom_configuration", return_value=custom_configuration), patch.object(manager, "_to_system_configuration", return_value=system_configuration), patch.object(manager, "_to_model_settings", return_value=[]), @@ -388,6 +434,7 @@ def test_get_configurations_reuses_cached_result_for_same_tenant(mocker: MockerF patch.object(manager, "_get_all_provider_model_settings", return_value={}), patch.object(manager, "_get_all_provider_load_balancing_configs", return_value={}), patch.object(manager, "_get_all_provider_model_credentials", return_value={}), + patch.object(manager, "_get_all_provider_credentials", return_value={}), patch.object(manager, "_to_custom_configuration", return_value=custom_configuration), patch.object(manager, "_to_system_configuration", return_value=system_configuration), patch.object(manager, "_to_model_settings", return_value=[]), @@ -424,6 +471,7 @@ def test_clear_configurations_cache_rebuilds_requested_tenant(mocker: MockerFixt patch.object(manager, "_get_all_provider_model_settings", return_value={}), patch.object(manager, "_get_all_provider_load_balancing_configs", return_value={}), patch.object(manager, "_get_all_provider_model_credentials", return_value={}), + patch.object(manager, "_get_all_provider_credentials", return_value={}), patch.object(manager, "_to_custom_configuration", return_value=custom_configuration), patch.object(manager, "_to_system_configuration", return_value=system_configuration), patch.object(manager, "_to_model_settings", return_value=[]), @@ -578,19 +626,162 @@ def test_get_all_providers_normalizes_provider_names_with_model_provider_id() -> assert list(result[str(ModelProviderID("langgenius/gemini/google"))]) == [gemini_provider] +def test_get_all_providers_attaches_active_credentials() -> None: + provider = Provider( + tenant_id="tenant-id", + provider_name="openai", + provider_type=ProviderType.CUSTOM, + is_valid=True, + credential_id="credential-id", + ) + provider.id = "provider-id" + credential = ProviderCredential( + tenant_id="tenant-id", + provider_name="openai", + credential_name="primary", + encrypted_config='{"api_key": "secret"}', + ) + credential.id = "credential-id" + session = Mock() + session.scalars.side_effect = [ + _FakeScalarResult([provider]), + _FakeScalarResult([credential]), + ] + + with ( + patch("core.provider_manager.session_factory.create_session", return_value=_build_session_context(session)), + ): + result = ProviderManager._get_all_providers("tenant-id") + + assert session.scalars.call_count == 2 + assert result[str(ModelProviderID("openai"))][0].credential_name == "primary" + assert result[str(ModelProviderID("openai"))][0].encrypted_config == '{"api_key": "secret"}' + + +def test_invalidate_configurations_cache_bumps_selected_source_version() -> None: + fake_redis = _FakeRedis() + + with patch("core.provider_manager.redis_client", fake_redis): + ProviderManager.invalidate_configurations_cache( + "tenant-id", + sources=(ProviderConfigurationCacheSource.PROVIDER_CREDENTIALS,), + ) + ProviderManager.invalidate_configurations_cache( + "tenant-id", + sources=(ProviderConfigurationCacheSource.PROVIDER_CREDENTIALS,), + ) + + assert fake_redis.store["provider_configurations:tenant:tenant-id:source:provider_credentials:version"] == "2" + assert fake_redis.expirations["provider_configurations:tenant:tenant-id:source:provider_credentials:version"] == 360 + assert "provider_configurations:tenant:tenant-id:source:provider_models:version" not in fake_redis.store + + +def test_provider_model_credentials_cache_returns_cache_entries() -> None: + fake_redis = _FakeRedis() + credential_record = SimpleNamespace( + id="credential-id", + provider_name="openai", + model_name="gpt-4", + model_type=ModelType.LLM, + credential_name="primary", + ) + session = Mock() + session.scalars.return_value = [credential_record] + + with ( + patch("core.provider_manager.redis_client", fake_redis), + patch("core.provider_manager.session_factory.create_session", return_value=_build_session_context(session)), + ): + first = ProviderManager._get_all_provider_model_credentials("tenant-id") + second = ProviderManager._get_all_provider_model_credentials("tenant-id") + + assert session.scalars.call_count == 1 + version_key = "provider_configurations:tenant:tenant-id:source:provider_model_credentials:version" + assert fake_redis.expirations[version_key] == 360 + assert first["openai"][0] is not credential_record + assert second["openai"][0].credential_name == "primary" + assert second["openai"][0].model_type == ModelType.LLM + + +def test_provider_configuration_cache_skips_write_when_version_changes_during_load() -> None: + fake_redis = _FakeRedis() + version_key = "provider_configurations:tenant:tenant-id:source:provider_model_credentials:version" + credential_record = SimpleNamespace( + id="credential-id", + provider_name="openai", + model_name="gpt-4", + model_type=ModelType.LLM, + credential_name="primary", + ) + session = Mock() + + def load_records(_stmt): + fake_redis.incr(version_key) + return [credential_record] + + session.scalars.side_effect = load_records + + with ( + patch("core.provider_manager.redis_client", fake_redis), + patch("core.provider_manager.session_factory.create_session", return_value=_build_session_context(session)), + ): + result = ProviderManager._get_all_provider_model_credentials("tenant-id") + + assert fake_redis.store[version_key] == "1" + assert "provider_configurations:tenant:tenant-id:source:provider_model_credentials:v:0" not in fake_redis.store + assert "provider_configurations:tenant:tenant-id:source:provider_model_credentials:v:1" not in fake_redis.store + assert result["openai"][0].credential_name == "primary" + + @pytest.mark.parametrize( "method_name", [ "_get_all_provider_models", "_get_all_provider_model_settings", "_get_all_provider_model_credentials", + "_get_all_provider_credentials", ], ) def test_provider_grouping_helpers_group_records_by_provider_name(method_name: str) -> None: + def build_record(provider_name: str, index: int): + match method_name: + case "_get_all_provider_models": + return SimpleNamespace( + id=f"model-{index}", + provider_name=provider_name, + model_name=f"model-{index}", + model_type=ModelType.LLM, + credential_id=None, + ) + case "_get_all_provider_model_settings": + return SimpleNamespace( + provider_name=provider_name, + model_name=f"model-{index}", + model_type=ModelType.LLM, + enabled=True, + load_balancing_enabled=False, + ) + case "_get_all_provider_model_credentials": + return SimpleNamespace( + id=f"model-credential-{index}", + provider_name=provider_name, + model_name=f"model-{index}", + model_type=ModelType.LLM, + credential_name=f"credential-{index}", + ) + case "_get_all_provider_credentials": + return SimpleNamespace( + id=f"credential-{index}", + provider_name=provider_name, + credential_name=f"credential-{index}", + ) + case _: + raise AssertionError(f"Unexpected method: {method_name}") + session = Mock() - openai_primary = SimpleNamespace(provider_name="openai") - openai_secondary = SimpleNamespace(provider_name="openai") - anthropic_record = SimpleNamespace(provider_name="anthropic") + openai_primary = build_record("openai", 1) + openai_secondary = build_record("openai", 2) + anthropic_record = build_record("anthropic", 3) session.scalars.return_value = [openai_primary, openai_secondary, anthropic_record] with ( @@ -598,14 +789,14 @@ def test_provider_grouping_helpers_group_records_by_provider_name(method_name: s ): result = getattr(ProviderManager, method_name)("tenant-id") - assert list(result["openai"]) == [openai_primary, openai_secondary] - assert list(result["anthropic"]) == [anthropic_record] + assert [record.provider_name for record in result["openai"]] == ["openai", "openai"] + assert [record.provider_name for record in result["anthropic"]] == ["anthropic"] def test_get_all_preferred_model_providers_returns_mapping_by_provider_name() -> None: session = Mock() - openai_preference = SimpleNamespace(provider_name="openai") - anthropic_preference = SimpleNamespace(provider_name="anthropic") + openai_preference = SimpleNamespace(provider_name="openai", preferred_provider_type=ProviderType.SYSTEM) + anthropic_preference = SimpleNamespace(provider_name="anthropic", preferred_provider_type=ProviderType.CUSTOM) session.scalars.return_value = [openai_preference, anthropic_preference] with ( @@ -613,10 +804,8 @@ def test_get_all_preferred_model_providers_returns_mapping_by_provider_name() -> ): result = ProviderManager._get_all_preferred_model_providers("tenant-id") - assert result == { - "openai": openai_preference, - "anthropic": anthropic_preference, - } + assert result["openai"].preferred_provider_type == ProviderType.SYSTEM + assert result["anthropic"].preferred_provider_type == ProviderType.CUSTOM def test_get_all_provider_load_balancing_configs_returns_empty_when_cached_flag_is_disabled() -> None: @@ -634,8 +823,30 @@ def test_get_all_provider_load_balancing_configs_returns_empty_when_cached_flag_ def test_get_all_provider_load_balancing_configs_populates_cache_and_groups_configs() -> None: session = Mock() - openai_config = SimpleNamespace(provider_name="openai") - anthropic_config = SimpleNamespace(provider_name="anthropic") + openai_config = SimpleNamespace( + id="lb-1", + tenant_id="tenant-id", + provider_name="openai", + model_name="gpt-4", + model_type=ModelType.LLM, + name="primary", + encrypted_config=None, + credential_id=None, + credential_source_type=None, + enabled=True, + ) + anthropic_config = SimpleNamespace( + id="lb-2", + tenant_id="tenant-id", + provider_name="anthropic", + model_name="claude", + model_type=ModelType.LLM, + name="primary", + encrypted_config=None, + credential_id=None, + credential_source_type=None, + enabled=True, + ) session.scalars.return_value = [openai_config, anthropic_config] with ( @@ -649,6 +860,6 @@ def test_get_all_provider_load_balancing_configs_populates_cache_and_groups_conf ): result = ProviderManager._get_all_provider_load_balancing_configs("tenant-id") - mock_setex.assert_called_once_with("tenant:tenant-id:model_load_balancing_enabled", 120, "True") - assert list(result["openai"]) == [openai_config] - assert list(result["anthropic"]) == [anthropic_config] + mock_setex.assert_any_call("tenant:tenant-id:model_load_balancing_enabled", 120, "True") + assert [record.provider_name for record in result["openai"]] == ["openai"] + assert [record.provider_name for record in result["anthropic"]] == ["anthropic"]