feat(api): cache workflow provider configurations (#37980)

This commit is contained in:
林玮 (Jade Lin) 2026-06-26 16:04:37 +08:00 committed by GitHub
parent 3f2ef24755
commit 1dbda1463e
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
5 changed files with 1004 additions and 97 deletions

View File

@ -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]:
"""

View File

@ -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.

View File

@ -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)

View File

@ -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:

View File

@ -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"]