From 39c8e5dbd5b36bb32d58980a1091f8aad9948abe Mon Sep 17 00:00:00 2001 From: Bakr Annour Date: Wed, 22 Jul 2026 18:36:00 +0200 Subject: [PATCH 01/17] removed ProviderGateway class from provider domain --- api/domain/provider/_providergateway.py | 20 -------------------- 1 file changed, 20 deletions(-) diff --git a/api/domain/provider/_providergateway.py b/api/domain/provider/_providergateway.py index f780e92c7..29da46672 100644 --- a/api/domain/provider/_providergateway.py +++ b/api/domain/provider/_providergateway.py @@ -1,27 +1,7 @@ -from abc import ABC, abstractmethod from dataclasses import dataclass -from api.domain.model.entities import ModelType as RouterType -from api.domain.model.errors import ModelNotFoundError -from api.domain.provider.entities import ProviderType -from api.domain.provider.errors import ProviderNotReachableError - @dataclass class ProviderCapabilities: max_context_length: int | None vector_size: int | None = None - - -class ProviderGateway(ABC): - @abstractmethod - async def get_capabilities( - self, - router_type: RouterType, - provider_type: ProviderType, - url: str, - key: str | None, - timeout: int, - model_name: str, - ) -> ProviderCapabilities | ModelNotFoundError | ProviderNotReachableError: - pass From 1151b79879198644e140c7e0334cdc19a5e3f93f Mon Sep 17 00:00:00 2001 From: Bakr Annour Date: Wed, 22 Jul 2026 18:37:23 +0200 Subject: [PATCH 02/17] refacto(model): remove ModelProviderGateway --- api/infrastructure/model/__init__.py | 3 - .../model/_modelprovidergateway.py | 99 ------------------- 2 files changed, 102 deletions(-) delete mode 100644 api/infrastructure/model/__init__.py delete mode 100644 api/infrastructure/model/_modelprovidergateway.py diff --git a/api/infrastructure/model/__init__.py b/api/infrastructure/model/__init__.py deleted file mode 100644 index 6597b06c7..000000000 --- a/api/infrastructure/model/__init__.py +++ /dev/null @@ -1,3 +0,0 @@ -from api.infrastructure.model._modelprovidergateway import ModelProviderGateway - -__all__ = ["ModelProviderGateway"] diff --git a/api/infrastructure/model/_modelprovidergateway.py b/api/infrastructure/model/_modelprovidergateway.py deleted file mode 100644 index b70b23c4e..000000000 --- a/api/infrastructure/model/_modelprovidergateway.py +++ /dev/null @@ -1,99 +0,0 @@ -import logging - -from api.domain.embeddings.entities import CreateEmbeddingsBody -from api.domain.model.entities import ModelType as RouterType -from api.domain.model.errors import ModelNotFoundError -from api.domain.provider import ProviderAdapterBuilder, ProviderCapabilities, ProviderClient, ProviderGateway -from api.domain.provider.entities import Provider, ProviderOriginalRequest, ProviderOriginalResponse, ProviderType -from api.domain.provider.errors import ProviderNotReachableError -from api.infrastructure.http.adapters.embeddings import EmbeddingsAdapter -from api.infrastructure.http.adapters.models import ModelsAdapter -from api.utils.variables import EndpointRoute - -logger = logging.getLogger(__name__) - - -class ModelProviderGateway(ProviderGateway): - def __init__(self, provider_client: ProviderClient, provider_adapter_builder: ProviderAdapterBuilder): - self.provider_adapter_builder = provider_adapter_builder - self.client = provider_client - - async def get_capabilities( - self, - router_type: RouterType, - provider_type: ProviderType, - url: str, - key: str | None, - timeout: int, - model_name: str, - ) -> ProviderCapabilities | ModelNotFoundError | ProviderNotReachableError: - provider = Provider( - id=0, - user_id=0, - router_id=0, - type=provider_type, - url=url, - key=key, - timeout=timeout, - model_name=model_name, - created=0, - updated=0, - ) - adapter = self.provider_adapter_builder.build(endpoint=EndpointRoute.MODELS, provider=provider) - - result = await self._get_max_context_length(adapter=adapter) - match result: - case ProviderNotReachableError() as error: - return error - case ModelNotFoundError() as error: - return error - case _: - max_context_length = result - - vector_size = None - if router_type == RouterType.TEXT_EMBEDDINGS_INFERENCE: - adapter = self.provider_adapter_builder.build(endpoint=EndpointRoute.EMBEDDINGS, provider=provider) - result = await self._get_vector_size(adapter=adapter) - match result: - case ProviderNotReachableError() as error: - return error - case _: - vector_size = result - - return ProviderCapabilities(max_context_length=max_context_length, vector_size=vector_size) - - async def _get_max_context_length(self, adapter: ModelsAdapter) -> int | None | ModelNotFoundError | ProviderNotReachableError: - original_request = ProviderOriginalRequest(endpoint=EndpointRoute.MODELS) - formatted_request = adapter.format_request(original_request=original_request) - response = await self.client.forward_request(provider=adapter.provider, formatted_request=formatted_request) - match response: - case ProviderOriginalResponse() as response: - pass - case error: - return ProviderNotReachableError(model_name=adapter.provider.model_name, status_code=error.status_code, detail=error.detail) - - formatted_response = adapter.format_response(original_response=response, original_request=original_request) - model_name = adapter.provider.model_name - model = next((model for model in formatted_response.data.data if model.id == model_name or model_name in model.aliases), None) - if model is None: - return ModelNotFoundError(name=model_name) - - return model.max_context_length - - async def _get_vector_size(self, adapter: EmbeddingsAdapter) -> int | ProviderNotReachableError: - original_request = ProviderOriginalRequest( - endpoint=EndpointRoute.EMBEDDINGS, - body=CreateEmbeddingsBody(model=adapter.provider.model_name, input="hello world"), - ) - formatted_request = adapter.format_request(original_request=original_request) - response = await self.client.forward_request(provider=adapter.provider, formatted_request=formatted_request) - match response: - case ProviderOriginalResponse() as response: - pass - case error: - return ProviderNotReachableError(model_name=adapter.provider.model_name, status_code=error.status_code, detail=error.detail) - - formatted_response = adapter.format_response(original_response=response, original_request=original_request) - vector_size = len(formatted_response.data.data[0].embedding) - - return vector_size From d82e9228e5c66beac2c7ab49730671cbe8e383e0 Mon Sep 17 00:00:00 2001 From: Bakr Annour Date: Thu, 23 Jul 2026 10:59:20 +0200 Subject: [PATCH 03/17] refacto(provider): add get_provider_capabilities use-case helper Extract provider capability-fetching (formerly ModelProviderGateway) into a shared module-level function in the use-case layer. Co-Authored-By: Claude Opus 4.8 --- api/use_cases/provider/__init__.py | 5 + .../provider/_getprovidercapabilities.py | 91 +++++++++++++++++++ 2 files changed, 96 insertions(+) create mode 100644 api/use_cases/provider/__init__.py create mode 100644 api/use_cases/provider/_getprovidercapabilities.py diff --git a/api/use_cases/provider/__init__.py b/api/use_cases/provider/__init__.py new file mode 100644 index 000000000..4ce07d5b2 --- /dev/null +++ b/api/use_cases/provider/__init__.py @@ -0,0 +1,5 @@ +from ._getprovidercapabilities import get_provider_capabilities + +__all__ = [ + "get_provider_capabilities", +] diff --git a/api/use_cases/provider/_getprovidercapabilities.py b/api/use_cases/provider/_getprovidercapabilities.py new file mode 100644 index 000000000..2e0ca32af --- /dev/null +++ b/api/use_cases/provider/_getprovidercapabilities.py @@ -0,0 +1,91 @@ +from api.domain.embeddings.entities import CreateEmbeddingsBody +from api.domain.model.entities import ModelType as RouterType +from api.domain.model.errors import ModelNotFoundError +from api.domain.provider import ProviderAdapter, ProviderAdapterBuilder, ProviderCapabilities, ProviderClient +from api.domain.provider.entities import Provider, ProviderOriginalRequest, ProviderOriginalResponse, ProviderType +from api.domain.provider.errors import ProviderNotReachableError +from api.utils.variables import EndpointRoute + + +async def get_provider_capabilities( + provider_client: ProviderClient, + provider_adapter_builder: ProviderAdapterBuilder, + router_type: RouterType, + provider_type: ProviderType, + url: str, + key: str | None, + timeout: int, + model_name: str, +) -> ProviderCapabilities | ModelNotFoundError | ProviderNotReachableError: + provider = Provider( + id=0, + user_id=0, + router_id=0, + type=provider_type, + url=url, + key=key, + timeout=timeout, + model_name=model_name, + created=0, + updated=0, + ) + adapter = provider_adapter_builder.build(endpoint=EndpointRoute.MODELS, provider=provider) + + result = await _get_max_context_length(provider_client=provider_client, adapter=adapter) + match result: + case ProviderNotReachableError() as error: + return error + case ModelNotFoundError() as error: + return error + case _: + max_context_length = result + + vector_size = None + if router_type == RouterType.TEXT_EMBEDDINGS_INFERENCE: + adapter = provider_adapter_builder.build(endpoint=EndpointRoute.EMBEDDINGS, provider=provider) + result = await _get_vector_size(provider_client=provider_client, adapter=adapter) + match result: + case ProviderNotReachableError() as error: + return error + case _: + vector_size = result + + return ProviderCapabilities(max_context_length=max_context_length, vector_size=vector_size) + + +async def _get_max_context_length(provider_client: ProviderClient, adapter: ProviderAdapter) -> int | None | ModelNotFoundError | ProviderNotReachableError: + original_request = ProviderOriginalRequest(endpoint=EndpointRoute.MODELS) + formatted_request = adapter.format_request(original_request=original_request) + response = await provider_client.forward_request(provider=adapter.provider, formatted_request=formatted_request) + match response: + case ProviderOriginalResponse() as response: + pass + case error: + return ProviderNotReachableError(model_name=adapter.provider.model_name, status_code=error.status_code, detail=error.detail) + + formatted_response = adapter.format_response(original_response=response, original_request=original_request) + model_name = adapter.provider.model_name + model = next((model for model in formatted_response.data.data if model.id == model_name or model_name in model.aliases), None) + if model is None: + return ModelNotFoundError(name=model_name) + + return model.max_context_length + + +async def _get_vector_size(provider_client: ProviderClient, adapter: ProviderAdapter) -> int | ProviderNotReachableError: + original_request = ProviderOriginalRequest( + endpoint=EndpointRoute.EMBEDDINGS, + body=CreateEmbeddingsBody(model=adapter.provider.model_name, input="hello world"), + ) + formatted_request = adapter.format_request(original_request=original_request) + response = await provider_client.forward_request(provider=adapter.provider, formatted_request=formatted_request) + match response: + case ProviderOriginalResponse() as response: + pass + case error: + return ProviderNotReachableError(model_name=adapter.provider.model_name, status_code=error.status_code, detail=error.detail) + + formatted_response = adapter.format_response(original_response=response, original_request=original_request) + vector_size = len(formatted_response.data.data[0].embedding) + + return vector_size From e2ddc69a6759696ae9efce3bb2d1bf98bd8a7e89 Mon Sep 17 00:00:00 2001 From: Bakr Annour Date: Thu, 23 Jul 2026 11:00:18 +0200 Subject: [PATCH 04/17] refacto(provider): delegate capability fetching to get_provider_capabilities Replace the removed provider_gateway dependency in CreateProviderUseCase and BootstrapModelsUseCase with calls to the shared get_provider_capabilities helper, injecting provider_client and provider_adapter_builder instead. Co-Authored-By: Claude Opus 4.8 --- .../admin/providers/_createproviderusecase.py | 12 ++++++++---- api/use_cases/models/_bootstrapmodelsusecase.py | 12 ++++++++---- 2 files changed, 16 insertions(+), 8 deletions(-) diff --git a/api/use_cases/admin/providers/_createproviderusecase.py b/api/use_cases/admin/providers/_createproviderusecase.py index bc8e2e09c..854618aa7 100644 --- a/api/use_cases/admin/providers/_createproviderusecase.py +++ b/api/use_cases/admin/providers/_createproviderusecase.py @@ -1,11 +1,12 @@ from dataclasses import dataclass from api.domain.model.errors import InconsistentModelMaxContextLengthError, InconsistentModelVectorSizeError -from api.domain.provider import ProviderGateway, ProviderRepository +from api.domain.provider import ProviderAdapterBuilder, ProviderClient, ProviderRepository from api.domain.provider.entities import BasicAuth, HostingZone, Metric, Provider, ProviderType from api.domain.provider.errors import InvalidProviderTypeError, ProviderAlreadyExistsError, ProviderNotReachableError from api.domain.router import RouterRepository from api.domain.router.errors import RouterNotFoundError +from api.use_cases.provider import get_provider_capabilities @dataclass @@ -42,10 +43,11 @@ class CreateProviderUseCaseSuccess: class CreateProviderUseCase: - def __init__(self, router_repository: RouterRepository, provider_repository: ProviderRepository, provider_gateway: ProviderGateway): + def __init__(self, router_repository: RouterRepository, provider_repository: ProviderRepository, provider_client: ProviderClient, provider_adapter_builder: ProviderAdapterBuilder): self.router_repository = router_repository self.provider_repository = provider_repository - self.provider_gateway = provider_gateway + self.provider_client = provider_client + self.provider_adapter_builder = provider_adapter_builder async def execute(self, command: CreateProviderCommand) -> CreateProviderUseCaseResult: router = await self.router_repository.get_router_by_id(router_id=command.router_id) @@ -55,7 +57,9 @@ async def execute(self, command: CreateProviderCommand) -> CreateProviderUseCase if not command.provider_type.is_compatible_with(router_type=router.type): return InvalidProviderTypeError(provider_type=command.provider_type.value, router_type=router.type.value) - result = await self.provider_gateway.get_capabilities( + result = await get_provider_capabilities( + provider_client=self.provider_client, + provider_adapter_builder=self.provider_adapter_builder, router_type=router.type, provider_type=command.provider_type, url=command.url, diff --git a/api/use_cases/models/_bootstrapmodelsusecase.py b/api/use_cases/models/_bootstrapmodelsusecase.py index 7b32b2e1c..2ba11a196 100644 --- a/api/use_cases/models/_bootstrapmodelsusecase.py +++ b/api/use_cases/models/_bootstrapmodelsusecase.py @@ -3,11 +3,12 @@ import logging from api.domain.model.errors import InconsistentModelMaxContextLengthError, InconsistentModelVectorSizeError, ModelNotFoundError -from api.domain.provider import ProviderGateway, ProviderRepository +from api.domain.provider import ProviderAdapterBuilder, ProviderClient, ProviderRepository from api.domain.provider.errors import ProviderAlreadyExistsError, ProviderNotReachableError from api.domain.router import RouterRepository from api.domain.router.errors import RouterNameAlreadyExistsError from api.schemas.core.configuration import Model as ModelConfiguration +from api.use_cases.provider import get_provider_capabilities logger = logging.getLogger(__name__) @@ -35,10 +36,11 @@ class BootstrapModelsUseCaseSkipped: class BootstrapModelsUseCase: - def __init__(self, router_repository: RouterRepository, provider_repository: ProviderRepository, provider_gateway: ProviderGateway): + def __init__(self, router_repository: RouterRepository, provider_repository: ProviderRepository, provider_client: ProviderClient, provider_adapter_builder: ProviderAdapterBuilder): self.router_repository = router_repository self.provider_repository = provider_repository - self.provider_gateway = provider_gateway + self.provider_client = provider_client + self.provider_adapter_builder = provider_adapter_builder async def execute( self, @@ -80,7 +82,9 @@ async def execute( ) for i, provider_to_create in enumerate(router_to_create.providers): - result = await self.provider_gateway.get_capabilities( + result = await get_provider_capabilities( + provider_client=self.provider_client, + provider_adapter_builder=self.provider_adapter_builder, router_type=router.type, provider_type=provider_to_create.type, url=provider_to_create.url, From 3f125d7a51e134b7b1c9922d99f569717eeceeb3 Mon Sep 17 00:00:00 2001 From: Bakr Annour Date: Thu, 23 Jul 2026 11:15:11 +0200 Subject: [PATCH 05/17] refacto(provider): drop ProviderGateway from wiring and domain exports Remove the deleted ModelProviderGateway/ProviderGateway references: inject provider_client and provider_adapter_builder into the create-provider factory, and drop ProviderGateway from the provider domain package exports. Co-Authored-By: Claude Opus 4.8 --- api/dependencies.py | 13 +++---------- api/domain/provider/__init__.py | 3 +-- 2 files changed, 4 insertions(+), 12 deletions(-) diff --git a/api/dependencies.py b/api/dependencies.py index 27a9bb3e8..6ff3f85a6 100644 --- a/api/dependencies.py +++ b/api/dependencies.py @@ -12,7 +12,6 @@ from api.domain.provider import ( ProviderAdapterBuilder, ProviderClient, - ProviderGateway, ProviderLoadBalancer, ProviderMetricsLogger, ProviderRepository, @@ -26,7 +25,6 @@ from api.infrastructure.fastapi.context import request_context from api.infrastructure.http import HttpProviderAdapterBuilder, HttpProviderClient from api.infrastructure.jwt import JwtKeyEncoder -from api.infrastructure.model import ModelProviderGateway from api.infrastructure.postgres import ( PostgresAuthenticatedUserQuery, PostgresKeyRepository, @@ -122,13 +120,6 @@ def _provider_client() -> ProviderClient: return HttpProviderClient() -def _provider_gateway( - provider_client: ProviderClient = Depends(_provider_client), - provider_adapter_builder: ProviderAdapterBuilder = Depends(_provider_adapter_builder), -) -> ProviderGateway: - return ModelProviderGateway(provider_client=provider_client, provider_adapter_builder=provider_adapter_builder) - - def _provider_load_balancer(redis_client: Redis = Depends(get_redis_client)) -> ProviderLoadBalancer: return RedisProviderLoadBalancer(redis_client=redis_client) @@ -363,11 +354,13 @@ def update_router_use_case_factory(postgres_session: AsyncSession = Depends(get_ def create_provider_use_case_factory( postgres_session: AsyncSession = Depends(get_postgres_session), provider_client: ProviderClient = Depends(_provider_client), + provider_adapter_builder: ProviderAdapterBuilder = Depends(_provider_adapter_builder), ) -> CreateProviderUseCase: return CreateProviderUseCase( router_repository=_router_repository(postgres_session), provider_repository=_provider_repository(postgres_session), - provider_gateway=_provider_gateway(provider_client=provider_client, provider_adapter_builder=HttpProviderAdapterBuilder()), + provider_client=provider_client, + provider_adapter_builder=provider_adapter_builder, ) diff --git a/api/domain/provider/__init__.py b/api/domain/provider/__init__.py index a976b5422..db8d70573 100644 --- a/api/domain/provider/__init__.py +++ b/api/domain/provider/__init__.py @@ -1,7 +1,7 @@ from api.domain.provider._provideradapter import ProviderAdapter from api.domain.provider._provideradapterbuilder import ProviderAdapterBuilder from api.domain.provider._providerclient import ProviderClient, ProviderClientResponse -from api.domain.provider._providergateway import ProviderCapabilities, ProviderGateway +from api.domain.provider._providergateway import ProviderCapabilities from api.domain.provider._providerloadbalancer import ProviderLoadBalancer from api.domain.provider._providermetricslogger import ProviderMetricsLogger from api.domain.provider._providerrepository import ProviderRepository @@ -12,7 +12,6 @@ "ProviderClient", "ProviderClientResponse", "ProviderCapabilities", - "ProviderGateway", "ProviderLoadBalancer", "ProviderMetricsLogger", "ProviderRepository", From c8461b942fc9898778889548967674580550cba2 Mon Sep 17 00:00:00 2001 From: Bakr Annour Date: Thu, 23 Jul 2026 11:17:43 +0200 Subject: [PATCH 06/17] refacto(provider): wire bootstrap_models without ModelProviderGateway --- api/utils/lifespan.py | 5 ++--- 1 file changed, 2 insertions(+), 3 deletions(-) diff --git a/api/utils/lifespan.py b/api/utils/lifespan.py index b100e43db..c87145fa6 100755 --- a/api/utils/lifespan.py +++ b/api/utils/lifespan.py @@ -19,7 +19,6 @@ from api.helpers.models import ModelRegistry from api.infrastructure.bcrypt import BcryptUserPasswordEncoder from api.infrastructure.http import HttpProviderAdapterBuilder, HttpProviderClient -from api.infrastructure.model import ModelProviderGateway from api.infrastructure.postgres import ( PostgresLimitRepository, PostgresPermissionRepository, @@ -124,12 +123,12 @@ async def bootstrap_models(configuration: Configuration, postgres_session: Async provider_repository = PostgresProviderRepository(postgres_session=postgres_session) provider_client = HttpProviderClient() provider_adapter_builder = HttpProviderAdapterBuilder() - provider_gateway = ModelProviderGateway(provider_client=provider_client, provider_adapter_builder=provider_adapter_builder) result = await BootstrapModelsUseCase( router_repository=router_repository, provider_repository=provider_repository, - provider_gateway=provider_gateway, + provider_client=provider_client, + provider_adapter_builder=provider_adapter_builder, ).execute(routers_to_create=configuration.models, bootstrap_admin_user_id=bootstrap_admin_user_id) match result: From 9af9ff0438756dd2c9f35dadad0975a86d269172 Mon Sep 17 00:00:00 2001 From: Bakr Annour Date: Thu, 23 Jul 2026 11:34:53 +0200 Subject: [PATCH 07/17] refacto(provider): move ProviderCapabilities into provider entities Relocate the ProviderCapabilities value object from the now-empty _providergateway module into entities.py (as a BaseModel, matching the other domain types) and delete _providergateway.py. Update importers to api.domain.provider.entities. Co-Authored-By: Claude Opus 4.8 --- api/domain/provider/__init__.py | 2 -- api/domain/provider/_providergateway.py | 7 ------- api/domain/provider/entities.py | 5 +++++ .../use_case/admin/providers/test_createproviderusecase.py | 3 +-- .../unit/use_case/admin/test_bootstrapmodelsusecase.py | 2 +- api/use_cases/provider/_getprovidercapabilities.py | 4 ++-- 6 files changed, 9 insertions(+), 14 deletions(-) delete mode 100644 api/domain/provider/_providergateway.py diff --git a/api/domain/provider/__init__.py b/api/domain/provider/__init__.py index db8d70573..8e80d707c 100644 --- a/api/domain/provider/__init__.py +++ b/api/domain/provider/__init__.py @@ -1,7 +1,6 @@ from api.domain.provider._provideradapter import ProviderAdapter from api.domain.provider._provideradapterbuilder import ProviderAdapterBuilder from api.domain.provider._providerclient import ProviderClient, ProviderClientResponse -from api.domain.provider._providergateway import ProviderCapabilities from api.domain.provider._providerloadbalancer import ProviderLoadBalancer from api.domain.provider._providermetricslogger import ProviderMetricsLogger from api.domain.provider._providerrepository import ProviderRepository @@ -11,7 +10,6 @@ "ProviderAdapterBuilder", "ProviderClient", "ProviderClientResponse", - "ProviderCapabilities", "ProviderLoadBalancer", "ProviderMetricsLogger", "ProviderRepository", diff --git a/api/domain/provider/_providergateway.py b/api/domain/provider/_providergateway.py deleted file mode 100644 index 29da46672..000000000 --- a/api/domain/provider/_providergateway.py +++ /dev/null @@ -1,7 +0,0 @@ -from dataclasses import dataclass - - -@dataclass -class ProviderCapabilities: - max_context_length: int | None - vector_size: int | None = None diff --git a/api/domain/provider/entities.py b/api/domain/provider/entities.py index da5774673..74800e5be 100644 --- a/api/domain/provider/entities.py +++ b/api/domain/provider/entities.py @@ -146,6 +146,11 @@ def is_compatible_with(self, router: Router) -> bool: return self.type.is_compatible_with(router.type) +class ProviderCapabilities(BaseModel): + max_context_length: int | None + vector_size: int | None = None + + class ProviderOriginalRequest(BaseModel): endpoint: Annotated[EndpointRoute, Field(description="The source endpoint (at the user side) of the request.")] body: Annotated[CreateEmbeddingsBody | CreateRerankBody | None, Field(default=None, description="The JSON body to use for the request.")] diff --git a/api/tests/unit/use_case/admin/providers/test_createproviderusecase.py b/api/tests/unit/use_case/admin/providers/test_createproviderusecase.py index e9c148e20..d5fedbc7e 100644 --- a/api/tests/unit/use_case/admin/providers/test_createproviderusecase.py +++ b/api/tests/unit/use_case/admin/providers/test_createproviderusecase.py @@ -4,8 +4,7 @@ from api.domain.model.entities import ModelType as RouterType from api.domain.model.errors import InconsistentModelMaxContextLengthError, InconsistentModelVectorSizeError -from api.domain.provider import ProviderCapabilities -from api.domain.provider.entities import HostingZone, ProviderType +from api.domain.provider.entities import HostingZone, ProviderCapabilities, ProviderType from api.domain.provider.errors import InvalidProviderTypeError, ProviderAlreadyExistsError, ProviderNotReachableError from api.domain.router.errors import RouterNotFoundError from api.tests.unit.use_case.factories import ProviderFactory, RouterFactory diff --git a/api/tests/unit/use_case/admin/test_bootstrapmodelsusecase.py b/api/tests/unit/use_case/admin/test_bootstrapmodelsusecase.py index cc2250aa4..155dbffdb 100644 --- a/api/tests/unit/use_case/admin/test_bootstrapmodelsusecase.py +++ b/api/tests/unit/use_case/admin/test_bootstrapmodelsusecase.py @@ -4,7 +4,7 @@ from api.domain.model.entities import ModelType as RouterType from api.domain.model.errors import InconsistentModelMaxContextLengthError, InconsistentModelVectorSizeError, ModelNotFoundError -from api.domain.provider import ProviderCapabilities +from api.domain.provider.entities import ProviderCapabilities from api.domain.provider.errors import ProviderAlreadyExistsError, ProviderNotReachableError from api.domain.router.errors import RouterNameAlreadyExistsError from api.tests.unit.use_case.factories import ( diff --git a/api/use_cases/provider/_getprovidercapabilities.py b/api/use_cases/provider/_getprovidercapabilities.py index 2e0ca32af..b7fe7c429 100644 --- a/api/use_cases/provider/_getprovidercapabilities.py +++ b/api/use_cases/provider/_getprovidercapabilities.py @@ -1,8 +1,8 @@ from api.domain.embeddings.entities import CreateEmbeddingsBody from api.domain.model.entities import ModelType as RouterType from api.domain.model.errors import ModelNotFoundError -from api.domain.provider import ProviderAdapter, ProviderAdapterBuilder, ProviderCapabilities, ProviderClient -from api.domain.provider.entities import Provider, ProviderOriginalRequest, ProviderOriginalResponse, ProviderType +from api.domain.provider import ProviderAdapter, ProviderAdapterBuilder, ProviderClient +from api.domain.provider.entities import Provider, ProviderCapabilities, ProviderOriginalRequest, ProviderOriginalResponse, ProviderType from api.domain.provider.errors import ProviderNotReachableError from api.utils.variables import EndpointRoute From 48e677fb08144ef52945867ceb0db36b069f6f3e Mon Sep 17 00:00:00 2001 From: Bakr Annour Date: Thu, 23 Jul 2026 14:23:42 +0200 Subject: [PATCH 08/17] refacto(provider): introduce ProviderCapabilitiesRepository instead of shared module in use-cases --- api/dependencies.py | 12 ++- api/domain/provider/__init__.py | 2 + .../_providercapabilitiesrepository.py | 97 +++++++++++++++++++ .../admin/providers/_createproviderusecase.py | 12 +-- .../models/_bootstrapmodelsusecase.py | 12 +-- api/use_cases/provider/__init__.py | 5 - .../provider/_getprovidercapabilities.py | 91 ----------------- api/utils/lifespan.py | 10 +- 8 files changed, 121 insertions(+), 120 deletions(-) create mode 100644 api/domain/provider/_providercapabilitiesrepository.py delete mode 100644 api/use_cases/provider/__init__.py delete mode 100644 api/use_cases/provider/_getprovidercapabilities.py diff --git a/api/dependencies.py b/api/dependencies.py index 6ff3f85a6..86b0be837 100644 --- a/api/dependencies.py +++ b/api/dependencies.py @@ -11,6 +11,7 @@ from api.domain.model import ModelEnvironmentalImpactsComputer, ModelTokenizer from api.domain.provider import ( ProviderAdapterBuilder, + ProviderCapabilitiesRepository, ProviderClient, ProviderLoadBalancer, ProviderMetricsLogger, @@ -160,6 +161,11 @@ def _permission_repository(session: AsyncSession) -> PermissionRepository: def _provider_repository(session: AsyncSession) -> ProviderRepository: return PostgresProviderRepository(postgres_session=session) +def _provider_capabilities_repository( + provider_client: ProviderClient = Depends(_provider_client), + provider_adapter_builder: ProviderAdapterBuilder = Depends(_provider_adapter_builder), +) -> ProviderCapabilitiesRepository: + return ProviderCapabilitiesRepository(provider_client=provider_client, provider_adapter_builder=provider_adapter_builder) # health use cases def get_health_models_use_case_factory( @@ -353,14 +359,12 @@ def update_router_use_case_factory(postgres_session: AsyncSession = Depends(get_ # provider use cases def create_provider_use_case_factory( postgres_session: AsyncSession = Depends(get_postgres_session), - provider_client: ProviderClient = Depends(_provider_client), - provider_adapter_builder: ProviderAdapterBuilder = Depends(_provider_adapter_builder), + provider_capabilities_repository: ProviderCapabilitiesRepository = Depends(_provider_capabilities_repository), ) -> CreateProviderUseCase: return CreateProviderUseCase( router_repository=_router_repository(postgres_session), provider_repository=_provider_repository(postgres_session), - provider_client=provider_client, - provider_adapter_builder=provider_adapter_builder, + provider_capabilities_repository=provider_capabilities_repository, ) diff --git a/api/domain/provider/__init__.py b/api/domain/provider/__init__.py index 8e80d707c..fe90bcb67 100644 --- a/api/domain/provider/__init__.py +++ b/api/domain/provider/__init__.py @@ -1,5 +1,6 @@ from api.domain.provider._provideradapter import ProviderAdapter from api.domain.provider._provideradapterbuilder import ProviderAdapterBuilder +from api.domain.provider._providercapabilitiesrepository import ProviderCapabilitiesRepository from api.domain.provider._providerclient import ProviderClient, ProviderClientResponse from api.domain.provider._providerloadbalancer import ProviderLoadBalancer from api.domain.provider._providermetricslogger import ProviderMetricsLogger @@ -13,4 +14,5 @@ "ProviderLoadBalancer", "ProviderMetricsLogger", "ProviderRepository", + "ProviderCapabilitiesRepository", ] diff --git a/api/domain/provider/_providercapabilitiesrepository.py b/api/domain/provider/_providercapabilitiesrepository.py new file mode 100644 index 000000000..1e56f66c1 --- /dev/null +++ b/api/domain/provider/_providercapabilitiesrepository.py @@ -0,0 +1,97 @@ +from api.domain.embeddings.entities import CreateEmbeddingsBody +from api.domain.model.entities import ModelType as RouterType +from api.domain.model.errors import ModelNotFoundError +from api.domain.provider import ProviderAdapter, ProviderAdapterBuilder, ProviderClient +from api.domain.provider.entities import Provider, ProviderCapabilities, ProviderOriginalRequest, ProviderOriginalResponse, ProviderType +from api.domain.provider.errors import ProviderNotReachableError +from api.utils.variables import EndpointRoute + + +class ProviderCapabilitiesRepository: + def __init__(self, provider_client: ProviderClient, provider_adapter_builder: ProviderAdapterBuilder): + self.provider_client = provider_client + self.provider_adapter_builder = provider_adapter_builder + + async def get_provider_capabilities( + self, + router_type: RouterType, + provider_type: ProviderType, + url: str, + key: str | None, + timeout: int, + model_name: str, + ) -> ProviderCapabilities | ModelNotFoundError | ProviderNotReachableError: + provider = Provider( + id=0, + user_id=0, + router_id=0, + type=provider_type, + url=url, + key=key, + timeout=timeout, + model_name=model_name, + created=0, + updated=0, + ) + adapter = self.provider_adapter_builder.build(endpoint=EndpointRoute.MODELS, provider=provider) + + result = await self._get_max_context_length(provider_client=self.provider_client, adapter=adapter) + match result: + case ProviderNotReachableError() as error: + return error + case ModelNotFoundError() as error: + return error + case _: + max_context_length = result + + vector_size = None + if router_type == RouterType.TEXT_EMBEDDINGS_INFERENCE: + adapter = self.provider_adapter_builder.build(endpoint=EndpointRoute.EMBEDDINGS, provider=provider) + result = await self._get_vector_size(provider_client=self.provider_client, adapter=adapter) + match result: + case ProviderNotReachableError() as error: + return error + case _: + vector_size = result + + return ProviderCapabilities(max_context_length=max_context_length, vector_size=vector_size) + + + @staticmethod + async def _get_max_context_length(provider_client: ProviderClient, adapter: ProviderAdapter) -> int | None | ModelNotFoundError | ProviderNotReachableError: + original_request = ProviderOriginalRequest(endpoint=EndpointRoute.MODELS) + formatted_request = adapter.format_request(original_request=original_request) + response = await provider_client.forward_request(provider=adapter.provider, formatted_request=formatted_request) + match response: + case ProviderOriginalResponse() as response: + pass + case error: + return ProviderNotReachableError(model_name=adapter.provider.model_name, status_code=error.status_code, detail=error.detail) + + formatted_response = adapter.format_response(original_response=response, original_request=original_request) + model_name = adapter.provider.model_name + model = next((model for model in formatted_response.data.data if model.id == model_name or model_name in model.aliases), None) + if model is None: + return ModelNotFoundError(name=model_name) + + return model.max_context_length + + + @staticmethod + async def _get_vector_size(provider_client: ProviderClient, adapter: ProviderAdapter) -> int | ProviderNotReachableError: + original_request = ProviderOriginalRequest( + endpoint=EndpointRoute.EMBEDDINGS, + body=CreateEmbeddingsBody(model=adapter.provider.model_name, input="hello world"), + ) + formatted_request = adapter.format_request(original_request=original_request) + response = await provider_client.forward_request(provider=adapter.provider, formatted_request=formatted_request) + match response: + case ProviderOriginalResponse() as response: + pass + case error: + return ProviderNotReachableError(model_name=adapter.provider.model_name, status_code=error.status_code, detail=error.detail) + + formatted_response = adapter.format_response(original_response=response, original_request=original_request) + vector_size = len(formatted_response.data.data[0].embedding) + + return vector_size diff --git a/api/use_cases/admin/providers/_createproviderusecase.py b/api/use_cases/admin/providers/_createproviderusecase.py index 854618aa7..e1e2ca6b6 100644 --- a/api/use_cases/admin/providers/_createproviderusecase.py +++ b/api/use_cases/admin/providers/_createproviderusecase.py @@ -1,12 +1,11 @@ from dataclasses import dataclass from api.domain.model.errors import InconsistentModelMaxContextLengthError, InconsistentModelVectorSizeError -from api.domain.provider import ProviderAdapterBuilder, ProviderClient, ProviderRepository +from api.domain.provider import ProviderCapabilitiesRepository, ProviderRepository from api.domain.provider.entities import BasicAuth, HostingZone, Metric, Provider, ProviderType from api.domain.provider.errors import InvalidProviderTypeError, ProviderAlreadyExistsError, ProviderNotReachableError from api.domain.router import RouterRepository from api.domain.router.errors import RouterNotFoundError -from api.use_cases.provider import get_provider_capabilities @dataclass @@ -43,11 +42,10 @@ class CreateProviderUseCaseSuccess: class CreateProviderUseCase: - def __init__(self, router_repository: RouterRepository, provider_repository: ProviderRepository, provider_client: ProviderClient, provider_adapter_builder: ProviderAdapterBuilder): + def __init__(self, router_repository: RouterRepository, provider_repository: ProviderRepository, provider_capabilities_repository: ProviderCapabilitiesRepository): self.router_repository = router_repository self.provider_repository = provider_repository - self.provider_client = provider_client - self.provider_adapter_builder = provider_adapter_builder + self.provider_capabilities_repository = provider_capabilities_repository async def execute(self, command: CreateProviderCommand) -> CreateProviderUseCaseResult: router = await self.router_repository.get_router_by_id(router_id=command.router_id) @@ -57,9 +55,7 @@ async def execute(self, command: CreateProviderCommand) -> CreateProviderUseCase if not command.provider_type.is_compatible_with(router_type=router.type): return InvalidProviderTypeError(provider_type=command.provider_type.value, router_type=router.type.value) - result = await get_provider_capabilities( - provider_client=self.provider_client, - provider_adapter_builder=self.provider_adapter_builder, + result = await self.provider_capabilities_repository.get_provider_capabilities( router_type=router.type, provider_type=command.provider_type, url=command.url, diff --git a/api/use_cases/models/_bootstrapmodelsusecase.py b/api/use_cases/models/_bootstrapmodelsusecase.py index 2ba11a196..53f73940f 100644 --- a/api/use_cases/models/_bootstrapmodelsusecase.py +++ b/api/use_cases/models/_bootstrapmodelsusecase.py @@ -3,12 +3,11 @@ import logging from api.domain.model.errors import InconsistentModelMaxContextLengthError, InconsistentModelVectorSizeError, ModelNotFoundError -from api.domain.provider import ProviderAdapterBuilder, ProviderClient, ProviderRepository +from api.domain.provider import ProviderCapabilitiesRepository, ProviderRepository from api.domain.provider.errors import ProviderAlreadyExistsError, ProviderNotReachableError from api.domain.router import RouterRepository from api.domain.router.errors import RouterNameAlreadyExistsError from api.schemas.core.configuration import Model as ModelConfiguration -from api.use_cases.provider import get_provider_capabilities logger = logging.getLogger(__name__) @@ -36,11 +35,10 @@ class BootstrapModelsUseCaseSkipped: class BootstrapModelsUseCase: - def __init__(self, router_repository: RouterRepository, provider_repository: ProviderRepository, provider_client: ProviderClient, provider_adapter_builder: ProviderAdapterBuilder): + def __init__(self, router_repository: RouterRepository, provider_repository: ProviderRepository, provider_capabilities_repository: ProviderCapabilitiesRepository): self.router_repository = router_repository self.provider_repository = provider_repository - self.provider_client = provider_client - self.provider_adapter_builder = provider_adapter_builder + self.provider_capabilities_repository = provider_capabilities_repository async def execute( self, @@ -82,9 +80,7 @@ async def execute( ) for i, provider_to_create in enumerate(router_to_create.providers): - result = await get_provider_capabilities( - provider_client=self.provider_client, - provider_adapter_builder=self.provider_adapter_builder, + result = await self.provider_capabilities_repository.get_provider_capabilities( router_type=router.type, provider_type=provider_to_create.type, url=provider_to_create.url, diff --git a/api/use_cases/provider/__init__.py b/api/use_cases/provider/__init__.py deleted file mode 100644 index 4ce07d5b2..000000000 --- a/api/use_cases/provider/__init__.py +++ /dev/null @@ -1,5 +0,0 @@ -from ._getprovidercapabilities import get_provider_capabilities - -__all__ = [ - "get_provider_capabilities", -] diff --git a/api/use_cases/provider/_getprovidercapabilities.py b/api/use_cases/provider/_getprovidercapabilities.py deleted file mode 100644 index b7fe7c429..000000000 --- a/api/use_cases/provider/_getprovidercapabilities.py +++ /dev/null @@ -1,91 +0,0 @@ -from api.domain.embeddings.entities import CreateEmbeddingsBody -from api.domain.model.entities import ModelType as RouterType -from api.domain.model.errors import ModelNotFoundError -from api.domain.provider import ProviderAdapter, ProviderAdapterBuilder, ProviderClient -from api.domain.provider.entities import Provider, ProviderCapabilities, ProviderOriginalRequest, ProviderOriginalResponse, ProviderType -from api.domain.provider.errors import ProviderNotReachableError -from api.utils.variables import EndpointRoute - - -async def get_provider_capabilities( - provider_client: ProviderClient, - provider_adapter_builder: ProviderAdapterBuilder, - router_type: RouterType, - provider_type: ProviderType, - url: str, - key: str | None, - timeout: int, - model_name: str, -) -> ProviderCapabilities | ModelNotFoundError | ProviderNotReachableError: - provider = Provider( - id=0, - user_id=0, - router_id=0, - type=provider_type, - url=url, - key=key, - timeout=timeout, - model_name=model_name, - created=0, - updated=0, - ) - adapter = provider_adapter_builder.build(endpoint=EndpointRoute.MODELS, provider=provider) - - result = await _get_max_context_length(provider_client=provider_client, adapter=adapter) - match result: - case ProviderNotReachableError() as error: - return error - case ModelNotFoundError() as error: - return error - case _: - max_context_length = result - - vector_size = None - if router_type == RouterType.TEXT_EMBEDDINGS_INFERENCE: - adapter = provider_adapter_builder.build(endpoint=EndpointRoute.EMBEDDINGS, provider=provider) - result = await _get_vector_size(provider_client=provider_client, adapter=adapter) - match result: - case ProviderNotReachableError() as error: - return error - case _: - vector_size = result - - return ProviderCapabilities(max_context_length=max_context_length, vector_size=vector_size) - - -async def _get_max_context_length(provider_client: ProviderClient, adapter: ProviderAdapter) -> int | None | ModelNotFoundError | ProviderNotReachableError: - original_request = ProviderOriginalRequest(endpoint=EndpointRoute.MODELS) - formatted_request = adapter.format_request(original_request=original_request) - response = await provider_client.forward_request(provider=adapter.provider, formatted_request=formatted_request) - match response: - case ProviderOriginalResponse() as response: - pass - case error: - return ProviderNotReachableError(model_name=adapter.provider.model_name, status_code=error.status_code, detail=error.detail) - - formatted_response = adapter.format_response(original_response=response, original_request=original_request) - model_name = adapter.provider.model_name - model = next((model for model in formatted_response.data.data if model.id == model_name or model_name in model.aliases), None) - if model is None: - return ModelNotFoundError(name=model_name) - - return model.max_context_length - - -async def _get_vector_size(provider_client: ProviderClient, adapter: ProviderAdapter) -> int | ProviderNotReachableError: - original_request = ProviderOriginalRequest( - endpoint=EndpointRoute.EMBEDDINGS, - body=CreateEmbeddingsBody(model=adapter.provider.model_name, input="hello world"), - ) - formatted_request = adapter.format_request(original_request=original_request) - response = await provider_client.forward_request(provider=adapter.provider, formatted_request=formatted_request) - match response: - case ProviderOriginalResponse() as response: - pass - case error: - return ProviderNotReachableError(model_name=adapter.provider.model_name, status_code=error.status_code, detail=error.detail) - - formatted_response = adapter.format_response(original_response=response, original_request=original_request) - vector_size = len(formatted_response.data.data[0].embedding) - - return vector_size diff --git a/api/utils/lifespan.py b/api/utils/lifespan.py index c87145fa6..bb20f0f72 100755 --- a/api/utils/lifespan.py +++ b/api/utils/lifespan.py @@ -9,6 +9,7 @@ from api.dependencies import get_postgres_session from api.domain.model.errors import InconsistentModelMaxContextLengthError, InconsistentModelVectorSizeError, ModelNotFoundError +from api.domain.provider import ProviderCapabilitiesRepository from api.domain.provider.errors import ProviderAlreadyExistsError, ProviderNotReachableError from api.domain.router.errors import RouterNameAlreadyExistsError from api.helpers._identityaccessmanager import IdentityAccessManager @@ -121,14 +122,15 @@ async def bootstrap_admin_role_and_user(configuration: Configuration, postgres_s async def bootstrap_models(configuration: Configuration, postgres_session: AsyncSession, bootstrap_admin_user_id: int) -> int: router_repository = PostgresRouterRepository(postgres_session=postgres_session, app_title=configuration.settings.app_title) provider_repository = PostgresProviderRepository(postgres_session=postgres_session) - provider_client = HttpProviderClient() - provider_adapter_builder = HttpProviderAdapterBuilder() + provider_capabilities_repository = ProviderCapabilitiesRepository( + provider_client=HttpProviderClient(), + provider_adapter_builder=HttpProviderAdapterBuilder(), + ) result = await BootstrapModelsUseCase( router_repository=router_repository, provider_repository=provider_repository, - provider_client=provider_client, - provider_adapter_builder=provider_adapter_builder, + provider_capabilities_repository=provider_capabilities_repository, ).execute(routers_to_create=configuration.models, bootstrap_admin_user_id=bootstrap_admin_user_id) match result: From 26cf6509f354020272ea8382763ffb286fc5047c Mon Sep 17 00:00:00 2001 From: Bakr Annour Date: Thu, 23 Jul 2026 15:11:08 +0200 Subject: [PATCH 09/17] refacto(provider): introduce ProviderCapabilitiesRepository Move the get_provider_capabilities use-case helper into a ProviderCapabilitiesRepository domain class and inject it where it's needed. - add ProviderCapabilitiesRepository under api/domain/provider, exposing get_provider_capabilities plus its context-length / vector-size helpers - wire it as a dependency and delegate from CreateProviderUseCase and BootstrapModelsUseCase, dropping their direct ProviderClient / ProviderAdapterBuilder dependencies - build it explicitly in the bootstrap_models lifespan hook - remove the now-unused api/use_cases/provider package Co-Authored-By: Claude Opus 4.8 --- .../providers/test_createproviderusecase.py | 45 ++++++++--------- .../admin/test_bootstrapmodelsusecase.py | 50 +++++++++---------- 2 files changed, 47 insertions(+), 48 deletions(-) diff --git a/api/tests/unit/use_case/admin/providers/test_createproviderusecase.py b/api/tests/unit/use_case/admin/providers/test_createproviderusecase.py index d5fedbc7e..8661c41d4 100644 --- a/api/tests/unit/use_case/admin/providers/test_createproviderusecase.py +++ b/api/tests/unit/use_case/admin/providers/test_createproviderusecase.py @@ -20,18 +20,17 @@ def router_repository(): def provider_repository(): return AsyncMock() - @pytest.fixture -def provider_gateway(): +def provider_capabilites_repository(): return AsyncMock() @pytest.fixture -def use_case(router_repository, provider_repository, provider_gateway): +def use_case(router_repository, provider_repository, provider_capabilites_repository): return CreateProviderUseCase( router_repository=router_repository, provider_repository=provider_repository, - provider_gateway=provider_gateway, + provider_capabilities_repository=provider_capabilites_repository, ) @@ -168,7 +167,7 @@ async def test_should_create_provider_when_router_has_a_different_provider( use_case, router_repository, provider_repository, - provider_gateway, + provider_capabilities_repository, sample_router_with_providers, sample_provider, default_command, @@ -176,7 +175,7 @@ async def test_should_create_provider_when_router_has_a_different_provider( # Arrange router_repository.get_router_by_id.return_value = sample_router_with_providers - provider_gateway.get_capabilities.return_value = ProviderCapabilities(max_context_length=4096, vector_size=None) + provider_capabilities_repository.get_provider_capabilities.return_value = ProviderCapabilities(max_context_length=4096, vector_size=None) provider_repository.create_provider.return_value = sample_provider # Act @@ -209,7 +208,7 @@ async def test_should_create_embedding_provider_when_vector_size_matches( use_case, router_repository, provider_repository, - provider_gateway, + provider_capabilities_repository, sample_embedding_router_with_providers, sample_provider, default_command, @@ -217,7 +216,7 @@ async def test_should_create_embedding_provider_when_vector_size_matches( # Arrange router_repository.get_router_by_id.return_value = sample_embedding_router_with_providers - provider_gateway.get_capabilities.return_value = ProviderCapabilities(max_context_length=512, vector_size=768) + provider_capabilities_repository.get_provider_capabilities.return_value = ProviderCapabilities(max_context_length=512, vector_size=768) provider_repository.create_provider.return_value = sample_provider # Act @@ -246,7 +245,7 @@ async def test_should_create_embedding_provider_when_vector_size_matches( @pytest.mark.asyncio async def test_should_return_router_not_found_error_when_router_does_not_exist( - self, use_case, router_repository, provider_repository, provider_gateway, default_command + self, use_case, router_repository, provider_repository, provider_capabilities_repository, default_command ): # Arrange @@ -258,7 +257,7 @@ async def test_should_return_router_not_found_error_when_router_does_not_exist( # Assert assert isinstance(result, RouterNotFoundError) assert result.id == 1 - provider_gateway.get_capabilities.assert_not_called() + provider_capabilities_repository.get_provider_capabilities.assert_not_called() provider_repository.create_provider.assert_not_called() @pytest.mark.asyncio @@ -272,7 +271,7 @@ async def test_should_create_provider_when_provider_type_is_compatible( use_case, router_repository, provider_repository, - provider_gateway, + provider_capabilities_repository, default_command, router_type, provider_type, @@ -281,7 +280,7 @@ async def test_should_create_provider_when_provider_type_is_compatible( capabilities = capabilities_for(router_type) provider = ProviderFactory(id=1, router_id=1, user_id=1, type=provider_type, url="https://example.com/", model_name="my-model") router_repository.get_router_by_id.return_value = RouterFactory(id=1, name="test-router", type=router_type, providers=0) - provider_gateway.get_capabilities.return_value = capabilities + provider_capabilities_repository.get_provider_capabilities.return_value = capabilities provider_repository.create_provider.return_value = provider command = with_provider_type(default_command, provider_type) @@ -291,7 +290,7 @@ async def test_should_create_provider_when_provider_type_is_compatible( # Assert assert isinstance(result, CreateProviderUseCaseSuccess) assert result.provider == provider - provider_gateway.get_capabilities.assert_called_once_with( + provider_capabilities_repository.get_provider_capabilities.assert_called_once_with( router_type=router_type, provider_type=provider_type, url="https://example.com/", @@ -328,7 +327,7 @@ async def test_should_return_invalid_provider_type_error_when_provider_type_is_n use_case, router_repository, provider_repository, - provider_gateway, + provider_capabilities_repository, default_command, router_type, provider_type, @@ -344,17 +343,17 @@ async def test_should_return_invalid_provider_type_error_when_provider_type_is_n assert isinstance(result, InvalidProviderTypeError) assert result.provider_type == provider_type.value assert result.router_type == router_type.value - provider_gateway.get_capabilities.assert_not_called() + provider_capabilities_repository.get_provider_capabilities.assert_not_called() provider_repository.create_provider.assert_not_called() @pytest.mark.asyncio async def test_should_return_provider_not_reachable_error_when_gateway_fails( - self, use_case, router_repository, provider_repository, provider_gateway, sample_router, default_command + self, use_case, router_repository, provider_repository, provider_capabilities_repository, sample_router, default_command ): # Arrange router_repository.get_router_by_id.return_value = sample_router - provider_gateway.get_capabilities.return_value = ProviderNotReachableError(model_name="my-model", status_code=500, detail="error_detail") + provider_capabilities_repository.get_provider_capabilities.return_value = ProviderNotReachableError(model_name="my-model", status_code=500, detail="error_detail") # Act result = await use_case.execute(default_command) @@ -368,12 +367,12 @@ async def test_should_return_provider_not_reachable_error_when_gateway_fails( @pytest.mark.asyncio async def test_should_return_inconsistent_max_context_length_error_when_mismatch( - self, use_case, router_repository, provider_repository, provider_gateway, sample_router_with_providers, default_command + self, use_case, router_repository, provider_repository, provider_capabilities_repository, sample_router_with_providers, default_command ): # Arrange router_repository.get_router_by_id.return_value = sample_router_with_providers - provider_gateway.get_capabilities.return_value = ProviderCapabilities(max_context_length=2048, vector_size=None) + provider_capabilities_repository.get_provider_capabilities.return_value = ProviderCapabilities(max_context_length=2048, vector_size=None) # Act result = await use_case.execute(default_command) @@ -390,14 +389,14 @@ async def test_should_return_inconsistent_vector_size_error_when_mismatch( use_case, router_repository, provider_repository, - provider_gateway, + provider_capabilities_repository, sample_embedding_router_with_providers, default_command, ): # Arrange router_repository.get_router_by_id.return_value = sample_embedding_router_with_providers - provider_gateway.get_capabilities.return_value = ProviderCapabilities(max_context_length=512, vector_size=384) + provider_capabilities_repository.get_provider_capabilities.return_value = ProviderCapabilities(max_context_length=512, vector_size=384) # Act result = await use_case.execute(with_provider_type(default_command, ProviderType.TEI)) @@ -410,12 +409,12 @@ async def test_should_return_inconsistent_vector_size_error_when_mismatch( @pytest.mark.asyncio async def test_should_return_provider_already_exists_error( - self, use_case, router_repository, provider_repository, provider_gateway, sample_router, default_command + self, use_case, router_repository, provider_repository, provider_capabilities_repository, sample_router, default_command ): # Arrange router_repository.get_router_by_id.return_value = sample_router - provider_gateway.get_capabilities.return_value = ProviderCapabilities(max_context_length=4096, vector_size=None) + provider_capabilities_repository.get_provider_capabilities.return_value = ProviderCapabilities(max_context_length=4096, vector_size=None) provider_repository.create_provider.return_value = ProviderAlreadyExistsError(model_name="my-model", url="https://example.com/", router_id=1) # Act diff --git a/api/tests/unit/use_case/admin/test_bootstrapmodelsusecase.py b/api/tests/unit/use_case/admin/test_bootstrapmodelsusecase.py index 155dbffdb..6cf2ce002 100644 --- a/api/tests/unit/use_case/admin/test_bootstrapmodelsusecase.py +++ b/api/tests/unit/use_case/admin/test_bootstrapmodelsusecase.py @@ -29,18 +29,18 @@ def provider_repository(): @pytest.fixture -def provider_gateway(): +def provider_capabilities_repository(): return AsyncMock() @pytest.fixture -def use_case(router_repository, provider_repository, provider_gateway): - return BootstrapModelsUseCase(router_repository=router_repository, provider_repository=provider_repository, provider_gateway=provider_gateway) +def use_case(router_repository, provider_repository, provider_capabilities_repository): + return BootstrapModelsUseCase(router_repository=router_repository, provider_repository=provider_repository, provider_capabilities_repository=provider_capabilities_repository) class TestBootstrapModelsUseCase: @pytest.mark.asyncio - async def test_skips_when_routers_already_exist(self, use_case, router_repository, provider_repository, provider_gateway): + async def test_skips_when_routers_already_exist(self, use_case, router_repository, provider_repository, provider_capabilities_repository): # Arrange existing_routers = [RouterFactory(id=1), RouterFactory(id=2)] router_repository.get_all_routers.return_value = existing_routers @@ -51,11 +51,11 @@ async def test_skips_when_routers_already_exist(self, use_case, router_repositor # Assert assert result == BootstrapModelsUseCaseSkipped(number_of_routers=2) router_repository.create_router.assert_not_awaited() - provider_gateway.get_capabilities.assert_not_awaited() + provider_capabilities_repository.get_provider_capabilities.assert_not_awaited() provider_repository.create_provider.assert_not_awaited() @pytest.mark.asyncio - async def test_successfully_creates_router_with_single_provider(self, use_case, router_repository, provider_repository, provider_gateway): + async def test_successfully_creates_router_with_single_provider(self, use_case, router_repository, provider_repository, provider_capabilities_repository): # Arrange model_provider = ModelProviderConfigurationFactory() model_configuration = ModelConfigurationFactory(providers=[model_provider]) @@ -64,7 +64,7 @@ async def test_successfully_creates_router_with_single_provider(self, use_case, router_repository.get_all_routers.return_value = [] router_repository.create_router.return_value = router - provider_gateway.get_capabilities.return_value = ProviderCapabilities(max_context_length=4096, vector_size=None) + provider_capabilities_repository.get_provider_capabilities.return_value = ProviderCapabilities(max_context_length=4096, vector_size=None) provider_repository.create_provider.return_value = provider # Act @@ -82,7 +82,7 @@ async def test_successfully_creates_router_with_single_provider(self, use_case, user_id=BOOTSTRAP_ADMIN_USER_ID, aliases=model_configuration.aliases, ) - provider_gateway.get_capabilities.assert_awaited_once_with( + provider_capabilities_repository.get_provider_capabilities.assert_awaited_once_with( router_type=router.type, provider_type=model_provider.type, url=model_provider.url, @@ -111,7 +111,7 @@ async def test_successfully_creates_router_with_single_provider(self, use_case, @pytest.mark.asyncio async def test_successfully_creates_multiple_routers_with_multiple_providers( - self, use_case, router_repository, provider_repository, provider_gateway + self, use_case, router_repository, provider_repository, provider_capabilities_repository ): # Arrange first_model = ModelConfigurationFactory( @@ -133,7 +133,7 @@ async def test_successfully_creates_multiple_routers_with_multiple_providers( router_repository.get_all_routers.return_value = [] router_repository.create_router.side_effect = [first_router, second_router] - provider_gateway.get_capabilities.side_effect = [ + provider_capabilities_repository.get_provider_capabilities.side_effect = [ ProviderCapabilities(max_context_length=4096, vector_size=None), ProviderCapabilities(max_context_length=4096, vector_size=None), ProviderCapabilities(max_context_length=512, vector_size=768), @@ -153,13 +153,13 @@ async def test_successfully_creates_multiple_routers_with_multiple_providers( # Assert assert result == BootstrapModelsUseCaseSuccess(number_of_routers=2) assert router_repository.create_router.await_count == 2 - assert provider_gateway.get_capabilities.await_count == 3 + assert provider_capabilities_repository.get_provider_capabilities.await_count == 3 assert provider_repository.create_provider.await_count == 3 router_repository.delete_all_routers.assert_not_awaited() @pytest.mark.asyncio async def test_returns_router_name_already_exists_error_when_duplicate_name( - self, use_case, router_repository, provider_repository, provider_gateway + self, use_case, router_repository, provider_repository, provider_capabilities_repository ): # Arrange routers_to_create = [ @@ -177,7 +177,7 @@ async def test_returns_router_name_already_exists_error_when_duplicate_name( # Assert assert result == RouterNameAlreadyExistsError(name="duplicate") router_repository.create_router.assert_not_awaited() - provider_gateway.get_capabilities.assert_not_awaited() + provider_capabilities_repository.get_provider_capabilities.assert_not_awaited() provider_repository.create_provider.assert_not_awaited() @pytest.mark.asyncio @@ -218,7 +218,7 @@ async def test_returns_provider_already_exists_error_when_duplicate_within_route use_case, router_repository, provider_repository, - provider_gateway, + provider_capabilities_repository, ): # Arrange model_configuration = ModelConfigurationFactory( @@ -237,17 +237,17 @@ async def test_returns_provider_already_exists_error_when_duplicate_within_route assert result.model_name == "model-a" assert result.url == "https://provider.com/" router_repository.create_router.assert_not_awaited() - provider_gateway.get_capabilities.assert_not_awaited() + provider_capabilities_repository.get_provider_capabilities.assert_not_awaited() provider_repository.create_provider.assert_not_awaited() @pytest.mark.asyncio - async def test_returns_provider_not_reachable_error_and_rolls_back(self, use_case, router_repository, provider_repository, provider_gateway): + async def test_returns_provider_not_reachable_error_and_rolls_back(self, use_case, router_repository, provider_repository, provider_capabilities_repository): # Arrange model_configuration = ModelConfigurationFactory() router = RouterFactory(id=1, name=model_configuration.name, type=RouterType.TEXT_GENERATION) router_repository.get_all_routers.return_value = [] router_repository.create_router.return_value = router - provider_gateway.get_capabilities.return_value = ProviderNotReachableError(model_name="my-model", status_code=500, detail="error_detail") + provider_capabilities_repository.get_provider_capabilities.return_value = ProviderNotReachableError(model_name="my-model", status_code=500, detail="error_detail") # Act result = await use_case.execute(routers_to_create=[model_configuration], bootstrap_admin_user_id=BOOTSTRAP_ADMIN_USER_ID) @@ -258,13 +258,13 @@ async def test_returns_provider_not_reachable_error_and_rolls_back(self, use_cas router_repository.delete_all_routers.assert_awaited_once() @pytest.mark.asyncio - async def test_returns_model_not_found_error_and_rolls_back(self, use_case, router_repository, provider_repository, provider_gateway): + async def test_returns_model_not_found_error_and_rolls_back(self, use_case, router_repository, provider_repository, provider_capabilities_repository): # Arrange model_configuration = ModelConfigurationFactory() router = RouterFactory(id=1, name=model_configuration.name, type=RouterType.TEXT_GENERATION) router_repository.get_all_routers.return_value = [] router_repository.create_router.return_value = router - provider_gateway.get_capabilities.return_value = ModelNotFoundError(name="my-model") + provider_capabilities_repository.get_provider_capabilities.return_value = ModelNotFoundError(name="my-model") # Act result = await use_case.execute(routers_to_create=[model_configuration], bootstrap_admin_user_id=BOOTSTRAP_ADMIN_USER_ID) @@ -280,7 +280,7 @@ async def test_returns_inconsistent_max_context_length_error_and_rolls_back( use_case, router_repository, provider_repository, - provider_gateway, + provider_capabilities_repository, ): # Arrange model_configuration = ModelConfigurationFactory( @@ -292,7 +292,7 @@ async def test_returns_inconsistent_max_context_length_error_and_rolls_back( router = RouterFactory(id=1, name=model_configuration.name, type=RouterType.TEXT_GENERATION, max_context_length=4096, vector_size=None) router_repository.get_all_routers.return_value = [] router_repository.create_router.return_value = router - provider_gateway.get_capabilities.side_effect = [ + provider_capabilities_repository.get_provider_capabilities.side_effect = [ ProviderCapabilities(max_context_length=4096, vector_size=None), ProviderCapabilities(max_context_length=2048, vector_size=None), ] @@ -309,7 +309,7 @@ async def test_returns_inconsistent_max_context_length_error_and_rolls_back( router_repository.delete_all_routers.assert_awaited_once() @pytest.mark.asyncio - async def test_returns_inconsistent_vector_size_error_and_rolls_back(self, use_case, router_repository, provider_repository, provider_gateway): + async def test_returns_inconsistent_vector_size_error_and_rolls_back(self, use_case, router_repository, provider_repository, provider_capabilities_repository): # Arrange model_configuration = ModelConfigurationFactory( type=RouterType.TEXT_EMBEDDINGS_INFERENCE, @@ -323,7 +323,7 @@ async def test_returns_inconsistent_vector_size_error_and_rolls_back(self, use_c ) router_repository.get_all_routers.return_value = [] router_repository.create_router.return_value = router - provider_gateway.get_capabilities.side_effect = [ + provider_capabilities_repository.get_provider_capabilities.side_effect = [ ProviderCapabilities(max_context_length=512, vector_size=768), ProviderCapabilities(max_context_length=512, vector_size=384), ] @@ -345,7 +345,7 @@ async def test_returns_inconsistent_vector_size_error_and_rolls_back(self, use_c router_repository.delete_all_routers.assert_awaited_once() @pytest.mark.asyncio - async def test_returns_success_with_no_routers_to_create(self, use_case, router_repository, provider_repository, provider_gateway): + async def test_returns_success_with_no_routers_to_create(self, use_case, router_repository, provider_repository, provider_capabilities_repository): # Arrange router_repository.get_all_routers.return_value = [] @@ -355,7 +355,7 @@ async def test_returns_success_with_no_routers_to_create(self, use_case, router_ # Assert assert result == BootstrapModelsUseCaseSuccess(number_of_routers=0) router_repository.create_router.assert_not_awaited() - provider_gateway.get_capabilities.assert_not_awaited() + provider_capabilities_repository.get_provider_capabilities.assert_not_awaited() provider_repository.create_provider.assert_not_awaited() From c15232752f96a4efb10826505820c66fe3b9581f Mon Sep 17 00:00:00 2001 From: Bakr Annour Date: Thu, 23 Jul 2026 15:26:36 +0200 Subject: [PATCH 10/17] test(provider): move capabilities integration test and fix import cycle - migrate the ModelProviderGateway integration test to ProviderCapabilitiesRepository under tests/integration/http - import ProviderClient / adapter classes from their submodules to break the circular import introduced by moving the repository into the provider package Co-Authored-By: Claude Opus 4.8 --- .../_providercapabilitiesrepository.py | 4 +++- .../test_providercapabilitiesrepository.py} | 19 +++++++++---------- 2 files changed, 12 insertions(+), 11 deletions(-) rename api/tests/integration/{models/test_modelprovidergateway.py => http/test_providercapabilitiesrepository.py} (83%) diff --git a/api/domain/provider/_providercapabilitiesrepository.py b/api/domain/provider/_providercapabilitiesrepository.py index 1e56f66c1..b04265a6e 100644 --- a/api/domain/provider/_providercapabilitiesrepository.py +++ b/api/domain/provider/_providercapabilitiesrepository.py @@ -1,7 +1,9 @@ from api.domain.embeddings.entities import CreateEmbeddingsBody from api.domain.model.entities import ModelType as RouterType from api.domain.model.errors import ModelNotFoundError -from api.domain.provider import ProviderAdapter, ProviderAdapterBuilder, ProviderClient +from api.domain.provider._provideradapter import ProviderAdapter +from api.domain.provider._provideradapterbuilder import ProviderAdapterBuilder +from api.domain.provider._providerclient import ProviderClient from api.domain.provider.entities import Provider, ProviderCapabilities, ProviderOriginalRequest, ProviderOriginalResponse, ProviderType from api.domain.provider.errors import ProviderNotReachableError from api.utils.variables import EndpointRoute diff --git a/api/tests/integration/models/test_modelprovidergateway.py b/api/tests/integration/http/test_providercapabilitiesrepository.py similarity index 83% rename from api/tests/integration/models/test_modelprovidergateway.py rename to api/tests/integration/http/test_providercapabilitiesrepository.py index 12bad0d6e..0597dc295 100644 --- a/api/tests/integration/models/test_modelprovidergateway.py +++ b/api/tests/integration/http/test_providercapabilitiesrepository.py @@ -5,10 +5,9 @@ import respx from api.domain.model.entities import ModelType as RouterType -from api.domain.provider import ProviderCapabilities -from api.domain.provider.entities import ProviderType +from api.domain.provider import ProviderCapabilitiesRepository +from api.domain.provider.entities import ProviderCapabilities, ProviderType from api.infrastructure.http import HttpProviderAdapterBuilder, HttpProviderClient -from api.infrastructure.model import ModelProviderGateway from api.tests.integration.factories.tei import TeiEmbeddingsResponseFactory, TeiModelsResponseFactory from api.tests.integration.factories.vllm import VllmModelsResponseFactory @@ -32,14 +31,14 @@ def _mock_embeddings_response(respx_mock, body: dict, status_code: int) -> None: @pytest.fixture -def gateway() -> ModelProviderGateway: - return ModelProviderGateway(provider_client=HttpProviderClient(), provider_adapter_builder=HttpProviderAdapterBuilder()) +def repository() -> ProviderCapabilitiesRepository: + return ProviderCapabilitiesRepository(provider_client=HttpProviderClient(), provider_adapter_builder=HttpProviderAdapterBuilder()) @pytest.mark.asyncio(loop_scope="session") -class TestModelProviderGateway: +class TestProviderCapabilitiesRepository: @respx.mock - async def test_get_capabilities_of_non_embeddings_providers(self, gateway: ModelProviderGateway): + async def test_get_capabilities_of_non_embeddings_providers(self, repository: ProviderCapabilitiesRepository): _mock_models_response( respx_mock=respx, provider_type=ProviderType.VLLM, @@ -47,7 +46,7 @@ async def test_get_capabilities_of_non_embeddings_providers(self, gateway: Model status_code=VllmModelsResponseFactory._status_code, ) - result = await gateway.get_capabilities( + result = await repository.get_provider_capabilities( router_type=RouterType.TEXT_GENERATION, provider_type=ProviderType.VLLM, url=DEFAULT_PROVIDER_URL, @@ -61,7 +60,7 @@ async def test_get_capabilities_of_non_embeddings_providers(self, gateway: Model @respx.mock async def test_get_capabilities_of_embeddings_providers( self, - gateway: ModelProviderGateway, + repository: ProviderCapabilitiesRepository, ): _mock_models_response( respx_mock=respx, @@ -75,7 +74,7 @@ async def test_get_capabilities_of_embeddings_providers( status_code=TeiEmbeddingsResponseFactory._status_code, ) - result = await gateway.get_capabilities( + result = await repository.get_provider_capabilities( router_type=RouterType.TEXT_EMBEDDINGS_INFERENCE, provider_type=ProviderType.TEI, url=DEFAULT_PROVIDER_URL, From 55cf5e014ce82ea4b199cd823af8e00adbe90c2b Mon Sep 17 00:00:00 2001 From: Bakr Annour Date: Thu, 23 Jul 2026 15:36:50 +0200 Subject: [PATCH 11/17] test(provider): migrate gateway unit tests to ProviderCapabilitiesRepository - move the orphaned ModelProviderGateway unit test to tests/unit/domain/provider, renaming symbols to ProviderCapabilitiesRepository - make _get_max_context_length / _get_vector_size instance methods that read self.provider_client, dropping the redundant argument threading Co-Authored-By: Claude Opus 4.8 --- .../_providercapabilitiesrepository.py | 14 ++-- api/tests/unit/domain/provider/__init__.py | 0 .../test_providercapabilitiesrepository.py} | 73 +++++++++---------- 3 files changed, 42 insertions(+), 45 deletions(-) create mode 100644 api/tests/unit/domain/provider/__init__.py rename api/tests/unit/{infrastructure/model/test_modelprovidergateway.py => domain/provider/test_providercapabilitiesrepository.py} (71%) diff --git a/api/domain/provider/_providercapabilitiesrepository.py b/api/domain/provider/_providercapabilitiesrepository.py index b04265a6e..5dee4ccbf 100644 --- a/api/domain/provider/_providercapabilitiesrepository.py +++ b/api/domain/provider/_providercapabilitiesrepository.py @@ -37,7 +37,7 @@ async def get_provider_capabilities( ) adapter = self.provider_adapter_builder.build(endpoint=EndpointRoute.MODELS, provider=provider) - result = await self._get_max_context_length(provider_client=self.provider_client, adapter=adapter) + result = await self._get_max_context_length(adapter=adapter) match result: case ProviderNotReachableError() as error: return error @@ -49,7 +49,7 @@ async def get_provider_capabilities( vector_size = None if router_type == RouterType.TEXT_EMBEDDINGS_INFERENCE: adapter = self.provider_adapter_builder.build(endpoint=EndpointRoute.EMBEDDINGS, provider=provider) - result = await self._get_vector_size(provider_client=self.provider_client, adapter=adapter) + result = await self._get_vector_size(adapter=adapter) match result: case ProviderNotReachableError() as error: return error @@ -59,11 +59,10 @@ async def get_provider_capabilities( return ProviderCapabilities(max_context_length=max_context_length, vector_size=vector_size) - @staticmethod - async def _get_max_context_length(provider_client: ProviderClient, adapter: ProviderAdapter) -> int | None | ModelNotFoundError | ProviderNotReachableError: + async def _get_max_context_length(self, adapter: ProviderAdapter) -> int | None | ModelNotFoundError | ProviderNotReachableError: original_request = ProviderOriginalRequest(endpoint=EndpointRoute.MODELS) formatted_request = adapter.format_request(original_request=original_request) - response = await provider_client.forward_request(provider=adapter.provider, formatted_request=formatted_request) + response = await self.provider_client.forward_request(provider=adapter.provider, formatted_request=formatted_request) match response: case ProviderOriginalResponse() as response: pass @@ -79,14 +78,13 @@ async def _get_max_context_length(provider_client: ProviderClient, adapter: Prov return model.max_context_length - @staticmethod - async def _get_vector_size(provider_client: ProviderClient, adapter: ProviderAdapter) -> int | ProviderNotReachableError: + async def _get_vector_size(self, adapter: ProviderAdapter) -> int | ProviderNotReachableError: original_request = ProviderOriginalRequest( endpoint=EndpointRoute.EMBEDDINGS, body=CreateEmbeddingsBody(model=adapter.provider.model_name, input="hello world"), ) formatted_request = adapter.format_request(original_request=original_request) - response = await provider_client.forward_request(provider=adapter.provider, formatted_request=formatted_request) + response = await self.provider_client.forward_request(provider=adapter.provider, formatted_request=formatted_request) match response: case ProviderOriginalResponse() as response: pass diff --git a/api/tests/unit/domain/provider/__init__.py b/api/tests/unit/domain/provider/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/api/tests/unit/infrastructure/model/test_modelprovidergateway.py b/api/tests/unit/domain/provider/test_providercapabilitiesrepository.py similarity index 71% rename from api/tests/unit/infrastructure/model/test_modelprovidergateway.py rename to api/tests/unit/domain/provider/test_providercapabilitiesrepository.py index d3f8a4874..e593a256d 100644 --- a/api/tests/unit/infrastructure/model/test_modelprovidergateway.py +++ b/api/tests/unit/domain/provider/test_providercapabilitiesrepository.py @@ -6,15 +6,14 @@ from api.domain.model.entities import Model, Models from api.domain.model.entities import ModelType as RouterType from api.domain.model.errors import ModelNotFoundError, StatusCodeModelError -from api.domain.provider import ProviderCapabilities -from api.domain.provider.entities import ProviderFormattedResponse, ProviderOriginalResponse, ProviderType +from api.domain.provider import ProviderCapabilitiesRepository +from api.domain.provider.entities import ProviderCapabilities, ProviderFormattedResponse, ProviderOriginalResponse, ProviderType from api.domain.provider.errors import ProviderNotReachableError from api.infrastructure.http import HttpProviderAdapterBuilder from api.infrastructure.http.adapters.embeddings import EmbeddingsAdapter from api.infrastructure.http.adapters.embeddings.tei import TeiEmbeddingsAdapter from api.infrastructure.http.adapters.models import ModelsAdapter from api.infrastructure.http.adapters.models.albert import AlbertModelsAdapter -from api.infrastructure.model._modelprovidergateway import ModelProviderGateway from api.tests.integration.factories.albert import AlbertModelResponseFactory, AlbertModelsResponseFactory from api.tests.integration.factories.tei import TeiEmbeddingsResponseFactory from api.tests.unit.use_case.factories import ProviderFactory @@ -36,8 +35,8 @@ def provider_adapter_builder() -> HttpProviderAdapterBuilder: @pytest.fixture -def gateway(provider_client: Mock, provider_adapter_builder: HttpProviderAdapterBuilder) -> ModelProviderGateway: - return ModelProviderGateway(provider_client=provider_client, provider_adapter_builder=provider_adapter_builder) +def repository(provider_client: Mock, provider_adapter_builder: HttpProviderAdapterBuilder) -> ProviderCapabilitiesRepository: + return ProviderCapabilitiesRepository(provider_client=provider_client, provider_adapter_builder=provider_adapter_builder) def provider_factory(provider_type: ProviderType = ProviderType.ALBERT, model_name: str = DEFAULT_MODEL_ID): @@ -52,13 +51,13 @@ def embeddings_adapter() -> TeiEmbeddingsAdapter: return TeiEmbeddingsAdapter(provider=provider_factory(provider_type=ProviderType.TEI)) -class TestModelProviderGateway: +class TestProviderCapabilitiesRepository: @pytest.mark.asyncio - async def test_should_get_capabilities_for_generation_router(self, gateway: ModelProviderGateway, mocker): - mocked_get_max_context_length = mocker.patch.object(ModelProviderGateway, "_get_max_context_length", AsyncMock(return_value=4096)) - mocked_get_vector_size = mocker.patch.object(ModelProviderGateway, "_get_vector_size", AsyncMock()) + async def test_should_get_capabilities_for_generation_router(self, repository: ProviderCapabilitiesRepository, mocker): + mocked_get_max_context_length = mocker.patch.object(ProviderCapabilitiesRepository, "_get_max_context_length", AsyncMock(return_value=4096)) + mocked_get_vector_size = mocker.patch.object(ProviderCapabilitiesRepository, "_get_vector_size", AsyncMock()) - result = await gateway.get_capabilities( + result = await repository.get_provider_capabilities( router_type=RouterType.TEXT_GENERATION, provider_type=ProviderType.ALBERT, url=DEFAULT_PROVIDER_URL, @@ -76,11 +75,11 @@ async def test_should_get_capabilities_for_generation_router(self, gateway: Mode mocked_get_vector_size.assert_not_called() @pytest.mark.asyncio - async def test_should_get_capabilities_for_embedding_router(self, gateway: ModelProviderGateway, mocker): - mocked_get_max_context_length = mocker.patch.object(ModelProviderGateway, "_get_max_context_length", AsyncMock(return_value=2048)) - mocked_get_vector_size = mocker.patch.object(ModelProviderGateway, "_get_vector_size", AsyncMock(return_value=3)) + async def test_should_get_capabilities_for_embedding_router(self, repository: ProviderCapabilitiesRepository, mocker): + mocked_get_max_context_length = mocker.patch.object(ProviderCapabilitiesRepository, "_get_max_context_length", AsyncMock(return_value=2048)) + mocked_get_vector_size = mocker.patch.object(ProviderCapabilitiesRepository, "_get_vector_size", AsyncMock(return_value=3)) - result = await gateway.get_capabilities( + result = await repository.get_provider_capabilities( router_type=RouterType.TEXT_EMBEDDINGS_INFERENCE, provider_type=ProviderType.TEI, url=DEFAULT_PROVIDER_URL, @@ -101,10 +100,10 @@ async def test_should_get_capabilities_for_embedding_router(self, gateway: Model "error", [ProviderNotReachableError(model_name=DEFAULT_MODEL_ID, status_code=500, detail="error_detail"), ModelNotFoundError(name=DEFAULT_MODEL_ID)], ) - async def test_should_return_max_context_error(self, gateway: ModelProviderGateway, error, mocker): - mocker.patch.object(ModelProviderGateway, "_get_max_context_length", AsyncMock(return_value=error)) + async def test_should_return_max_context_error(self, repository: ProviderCapabilitiesRepository, error, mocker): + mocker.patch.object(ProviderCapabilitiesRepository, "_get_max_context_length", AsyncMock(return_value=error)) - result = await gateway.get_capabilities( + result = await repository.get_provider_capabilities( router_type=RouterType.TEXT_GENERATION, provider_type=ProviderType.ALBERT, url=DEFAULT_PROVIDER_URL, @@ -116,12 +115,12 @@ async def test_should_return_max_context_error(self, gateway: ModelProviderGatew assert result == error @pytest.mark.asyncio - async def test_should_return_vector_size_error(self, gateway: ModelProviderGateway, mocker): + async def test_should_return_vector_size_error(self, repository: ProviderCapabilitiesRepository, mocker): error = ProviderNotReachableError(model_name=DEFAULT_MODEL_ID, status_code=500, detail="error_detail") - mocker.patch.object(ModelProviderGateway, "_get_max_context_length", AsyncMock(return_value=4096)) - mocker.patch.object(ModelProviderGateway, "_get_vector_size", AsyncMock(return_value=error)) + mocker.patch.object(ProviderCapabilitiesRepository, "_get_max_context_length", AsyncMock(return_value=4096)) + mocker.patch.object(ProviderCapabilitiesRepository, "_get_vector_size", AsyncMock(return_value=error)) - result = await gateway.get_capabilities( + result = await repository.get_provider_capabilities( router_type=RouterType.TEXT_EMBEDDINGS_INFERENCE, provider_type=ProviderType.TEI, url=DEFAULT_PROVIDER_URL, @@ -133,14 +132,14 @@ async def test_should_return_vector_size_error(self, gateway: ModelProviderGatew assert result == error @pytest.mark.asyncio - async def test_should_get_max_context_length_when_model_id_is_found(self, gateway: ModelProviderGateway, provider_client: Mock): + async def test_should_get_max_context_length_when_model_id_is_found(self, repository: ProviderCapabilitiesRepository, provider_client: Mock): body = AlbertModelsResponseFactory( count=2, data=[AlbertModelResponseFactory(model=DEFAULT_MODEL_ID, aliases=["test-model-alias"], max_context_length=10)], ) provider_client.forward_request.return_value = ProviderOriginalResponse(data=body) - result = await gateway._get_max_context_length(adapter=models_adapter()) + result = await repository._get_max_context_length(adapter=models_adapter()) assert result == 10 provider_client.forward_request.assert_awaited_once() @@ -149,7 +148,7 @@ async def test_should_get_max_context_length_when_model_id_is_found(self, gatewa assert formatted_request.url == f"{DEFAULT_PROVIDER_URL}/v1/models" @pytest.mark.asyncio - async def test_should_get_max_context_length_when_model_alias_is_found(self, gateway: ModelProviderGateway, provider_client: Mock): + async def test_should_get_max_context_length_when_model_alias_is_found(self, repository: ProviderCapabilitiesRepository, provider_client: Mock): adapter = Mock() adapter.provider = provider_factory(model_name="model-alias") adapter.format_request.return_value = Mock() @@ -177,13 +176,13 @@ async def test_should_get_max_context_length_when_model_alias_is_found(self, gat ) provider_client.forward_request.return_value = ProviderOriginalResponse(data={}) - result = await gateway._get_max_context_length(adapter=adapter) + result = await repository._get_max_context_length(adapter=adapter) assert result == 10 @pytest.mark.asyncio async def test_should_return_the_first_model_max_context_length_when_several_models_with_the_same_name_are_found( - self, gateway: ModelProviderGateway, provider_client: Mock + self, repository: ProviderCapabilitiesRepository, provider_client: Mock ): body = AlbertModelsResponseFactory( data=[ @@ -193,41 +192,41 @@ async def test_should_return_the_first_model_max_context_length_when_several_mod ) provider_client.forward_request.return_value = ProviderOriginalResponse(data=body) - result = await gateway._get_max_context_length(adapter=models_adapter()) + result = await repository._get_max_context_length(adapter=models_adapter()) assert result == 10 @pytest.mark.asyncio - async def test_should_return_model_not_found_when_models_response_is_empty(self, gateway: ModelProviderGateway, provider_client: Mock): + async def test_should_return_model_not_found_when_models_response_is_empty(self, repository: ProviderCapabilitiesRepository, provider_client: Mock): provider_client.forward_request.return_value = ProviderOriginalResponse(data=AlbertModelsResponseFactory(data=[])) - result = await gateway._get_max_context_length(adapter=models_adapter()) + result = await repository._get_max_context_length(adapter=models_adapter()) assert result == ModelNotFoundError(name=DEFAULT_MODEL_ID) @pytest.mark.asyncio - async def test_should_return_model_not_found_when_model_is_missing_in_models_response(self, gateway: ModelProviderGateway, provider_client: Mock): + async def test_should_return_model_not_found_when_model_is_missing_in_models_response(self, repository: ProviderCapabilitiesRepository, provider_client: Mock): provider_client.forward_request.return_value = ProviderOriginalResponse(data=AlbertModelsResponseFactory(data=[AlbertModelResponseFactory()])) - result = await gateway._get_max_context_length(adapter=models_adapter()) + result = await repository._get_max_context_length(adapter=models_adapter()) assert result == ModelNotFoundError(name=DEFAULT_MODEL_ID) @pytest.mark.asyncio - async def test_should_return_provider_not_reachable_when_getting_max_context_fails(self, gateway: ModelProviderGateway, provider_client: Mock): + async def test_should_return_provider_not_reachable_when_getting_max_context_fails(self, repository: ProviderCapabilitiesRepository, provider_client: Mock): provider_client.forward_request.return_value = StatusCodeModelError(status_code=500, detail="boom") - result = await gateway._get_max_context_length(adapter=models_adapter()) + result = await repository._get_max_context_length(adapter=models_adapter()) assert result == ProviderNotReachableError(model_name=DEFAULT_MODEL_ID, status_code=500, detail="boom") @pytest.mark.asyncio - async def test_should_get_vector_size(self, gateway: ModelProviderGateway, provider_client: Mock): + async def test_should_get_vector_size(self, repository: ProviderCapabilitiesRepository, provider_client: Mock): provider_client.forward_request.return_value = ProviderOriginalResponse( data=TeiEmbeddingsResponseFactory(dimensions=3, model_id=DEFAULT_MODEL_ID) ) - result = await gateway._get_vector_size(adapter=embeddings_adapter()) + result = await repository._get_vector_size(adapter=embeddings_adapter()) assert result == 3 provider_client.forward_request.assert_awaited_once() @@ -237,9 +236,9 @@ async def test_should_get_vector_size(self, gateway: ModelProviderGateway, provi assert formatted_request.body["model"] == DEFAULT_MODEL_ID @pytest.mark.asyncio - async def test_should_return_provider_not_reachable_when_getting_vector_size_fails(self, gateway: ModelProviderGateway, provider_client: Mock): + async def test_should_return_provider_not_reachable_when_getting_vector_size_fails(self, repository: ProviderCapabilitiesRepository, provider_client: Mock): provider_client.forward_request.return_value = StatusCodeModelError(status_code=500, detail="boom") - result = await gateway._get_vector_size(adapter=embeddings_adapter()) + result = await repository._get_vector_size(adapter=embeddings_adapter()) assert result == ProviderNotReachableError(model_name=DEFAULT_MODEL_ID, status_code=500, detail="boom") From 3303e223ab05de66faa290e5803741550ec31272 Mon Sep 17 00:00:00 2001 From: Bakr Annour Date: Thu, 23 Jul 2026 16:02:56 +0200 Subject: [PATCH 12/17] fix(provider): handle ModelNotFoundError in create provider --- .../fastapi/endpoints/admin/providers.py | 7 +++++- .../providers/test_createproviderusecase.py | 25 ++++++++++++++++--- .../admin/providers/_createproviderusecase.py | 5 +++- 3 files changed, 31 insertions(+), 6 deletions(-) diff --git a/api/infrastructure/fastapi/endpoints/admin/providers.py b/api/infrastructure/fastapi/endpoints/admin/providers.py index aac3d5e7e..69a0ac79a 100644 --- a/api/infrastructure/fastapi/endpoints/admin/providers.py +++ b/api/infrastructure/fastapi/endpoints/admin/providers.py @@ -12,7 +12,7 @@ update_provider_use_case_factory, ) from api.domain import SortOrder -from api.domain.model.errors import InconsistentModelMaxContextLengthError, InconsistentModelVectorSizeError +from api.domain.model.errors import InconsistentModelMaxContextLengthError, InconsistentModelVectorSizeError, ModelNotFoundError from api.domain.provider.entities import ProviderSortField from api.domain.provider.errors import InvalidProviderTypeError, ProviderAlreadyExistsError, ProviderNotFoundError, ProviderNotReachableError from api.domain.router.errors import RouterNotFoundError @@ -25,6 +25,7 @@ InconsistentModelVectorSizeHTTPException, InternalServerHTTPException, InvalidProviderTypeHTTPException, + ModelNotFoundHTTPException, NotAdminUserHTTPException, ProviderAlreadyExistsHTTPException, ProviderNotFoundHTTPException, @@ -70,6 +71,7 @@ InconsistentModelVectorSizeHTTPException, InvalidProviderTypeHTTPException, ProviderNotReachableHTTPException, + ModelNotFoundHTTPException, ProviderAlreadyExistsHTTPException, RouterNotFoundHTTPException, NotAdminUserHTTPException, @@ -127,6 +129,9 @@ async def create_provider( case ProviderNotReachableError() as error: raise ProviderNotReachableHTTPException(name=error.model_name, status_code=error.status_code, detail=error.detail) + case ModelNotFoundError(name=model_name): + raise ModelNotFoundHTTPException(name=model_name) + case ProviderAlreadyExistsError(model_name=model_name, url=url, router_id=router_id): raise ProviderAlreadyExistsHTTPException(model_name=model_name, url=url, router_id=router_id) diff --git a/api/tests/unit/use_case/admin/providers/test_createproviderusecase.py b/api/tests/unit/use_case/admin/providers/test_createproviderusecase.py index 8661c41d4..f75775804 100644 --- a/api/tests/unit/use_case/admin/providers/test_createproviderusecase.py +++ b/api/tests/unit/use_case/admin/providers/test_createproviderusecase.py @@ -3,7 +3,7 @@ import pytest from api.domain.model.entities import ModelType as RouterType -from api.domain.model.errors import InconsistentModelMaxContextLengthError, InconsistentModelVectorSizeError +from api.domain.model.errors import InconsistentModelMaxContextLengthError, InconsistentModelVectorSizeError, ModelNotFoundError from api.domain.provider.entities import HostingZone, ProviderCapabilities, ProviderType from api.domain.provider.errors import InvalidProviderTypeError, ProviderAlreadyExistsError, ProviderNotReachableError from api.domain.router.errors import RouterNotFoundError @@ -21,16 +21,16 @@ def provider_repository(): return AsyncMock() @pytest.fixture -def provider_capabilites_repository(): +def provider_capabilities_repository(): return AsyncMock() @pytest.fixture -def use_case(router_repository, provider_repository, provider_capabilites_repository): +def use_case(router_repository, provider_repository, provider_capabilities_repository): return CreateProviderUseCase( router_repository=router_repository, provider_repository=provider_repository, - provider_capabilities_repository=provider_capabilites_repository, + provider_capabilities_repository=provider_capabilities_repository, ) @@ -365,6 +365,23 @@ async def test_should_return_provider_not_reachable_error_when_gateway_fails( assert result.detail == "error_detail" provider_repository.create_provider.assert_not_called() + @pytest.mark.asyncio + async def test_should_return_model_not_found_error_when_model_is_missing( + self, use_case, router_repository, provider_repository, provider_capabilities_repository, sample_router, default_command + ): + # Arrange + + router_repository.get_router_by_id.return_value = sample_router + provider_capabilities_repository.get_provider_capabilities.return_value = ModelNotFoundError(name="my-model") + + # Act + result = await use_case.execute(default_command) + + # Assert + assert isinstance(result, ModelNotFoundError) + assert result.name == "my-model" + provider_repository.create_provider.assert_not_called() + @pytest.mark.asyncio async def test_should_return_inconsistent_max_context_length_error_when_mismatch( self, use_case, router_repository, provider_repository, provider_capabilities_repository, sample_router_with_providers, default_command diff --git a/api/use_cases/admin/providers/_createproviderusecase.py b/api/use_cases/admin/providers/_createproviderusecase.py index e1e2ca6b6..ff1d66aec 100644 --- a/api/use_cases/admin/providers/_createproviderusecase.py +++ b/api/use_cases/admin/providers/_createproviderusecase.py @@ -1,6 +1,6 @@ from dataclasses import dataclass -from api.domain.model.errors import InconsistentModelMaxContextLengthError, InconsistentModelVectorSizeError +from api.domain.model.errors import InconsistentModelMaxContextLengthError, InconsistentModelVectorSizeError, ModelNotFoundError from api.domain.provider import ProviderCapabilitiesRepository, ProviderRepository from api.domain.provider.entities import BasicAuth, HostingZone, Metric, Provider, ProviderType from api.domain.provider.errors import InvalidProviderTypeError, ProviderAlreadyExistsError, ProviderNotReachableError @@ -34,6 +34,7 @@ class CreateProviderUseCaseSuccess: CreateProviderUseCaseSuccess | InvalidProviderTypeError | ProviderNotReachableError + | ModelNotFoundError | InconsistentModelMaxContextLengthError | InconsistentModelVectorSizeError | RouterNotFoundError @@ -66,6 +67,8 @@ async def execute(self, command: CreateProviderCommand) -> CreateProviderUseCase match result: case ProviderNotReachableError() as error: return error + case ModelNotFoundError() as error: + return error case provider_capabilities: pass From b31b5eb4deffd54c53aea68489ffa2d1d5d578b3 Mon Sep 17 00:00:00 2001 From: Bakr Annour Date: Thu, 23 Jul 2026 16:19:42 +0200 Subject: [PATCH 13/17] Reformatting dependencies.py --- api/dependencies.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/api/dependencies.py b/api/dependencies.py index 86b0be837..15713a6d1 100644 --- a/api/dependencies.py +++ b/api/dependencies.py @@ -161,12 +161,14 @@ def _permission_repository(session: AsyncSession) -> PermissionRepository: def _provider_repository(session: AsyncSession) -> ProviderRepository: return PostgresProviderRepository(postgres_session=session) + def _provider_capabilities_repository( provider_client: ProviderClient = Depends(_provider_client), provider_adapter_builder: ProviderAdapterBuilder = Depends(_provider_adapter_builder), ) -> ProviderCapabilitiesRepository: return ProviderCapabilitiesRepository(provider_client=provider_client, provider_adapter_builder=provider_adapter_builder) + # health use cases def get_health_models_use_case_factory( postgres_session: AsyncSession = Depends(get_postgres_session), From 9538b1b73b920ebc08ac1485bc9a8f4dde73cb62 Mon Sep 17 00:00:00 2001 From: Bakr Annour Date: Thu, 23 Jul 2026 16:21:47 +0200 Subject: [PATCH 14/17] style(provider): apply ruff-format --- .../_providercapabilitiesrepository.py | 2 -- .../test_providercapabilitiesrepository.py | 16 +++++++--- .../providers/test_createproviderusecase.py | 5 +++- .../admin/test_bootstrapmodelsusecase.py | 30 ++++++++++++++----- .../admin/providers/_createproviderusecase.py | 7 ++++- .../models/_bootstrapmodelsusecase.py | 7 ++++- 6 files changed, 51 insertions(+), 16 deletions(-) diff --git a/api/domain/provider/_providercapabilitiesrepository.py b/api/domain/provider/_providercapabilitiesrepository.py index 5dee4ccbf..69e8f67a5 100644 --- a/api/domain/provider/_providercapabilitiesrepository.py +++ b/api/domain/provider/_providercapabilitiesrepository.py @@ -58,7 +58,6 @@ async def get_provider_capabilities( return ProviderCapabilities(max_context_length=max_context_length, vector_size=vector_size) - async def _get_max_context_length(self, adapter: ProviderAdapter) -> int | None | ModelNotFoundError | ProviderNotReachableError: original_request = ProviderOriginalRequest(endpoint=EndpointRoute.MODELS) formatted_request = adapter.format_request(original_request=original_request) @@ -77,7 +76,6 @@ async def _get_max_context_length(self, adapter: ProviderAdapter) -> int | None return model.max_context_length - async def _get_vector_size(self, adapter: ProviderAdapter) -> int | ProviderNotReachableError: original_request = ProviderOriginalRequest( endpoint=EndpointRoute.EMBEDDINGS, diff --git a/api/tests/unit/domain/provider/test_providercapabilitiesrepository.py b/api/tests/unit/domain/provider/test_providercapabilitiesrepository.py index e593a256d..8f9998c56 100644 --- a/api/tests/unit/domain/provider/test_providercapabilitiesrepository.py +++ b/api/tests/unit/domain/provider/test_providercapabilitiesrepository.py @@ -197,7 +197,9 @@ async def test_should_return_the_first_model_max_context_length_when_several_mod assert result == 10 @pytest.mark.asyncio - async def test_should_return_model_not_found_when_models_response_is_empty(self, repository: ProviderCapabilitiesRepository, provider_client: Mock): + async def test_should_return_model_not_found_when_models_response_is_empty( + self, repository: ProviderCapabilitiesRepository, provider_client: Mock + ): provider_client.forward_request.return_value = ProviderOriginalResponse(data=AlbertModelsResponseFactory(data=[])) result = await repository._get_max_context_length(adapter=models_adapter()) @@ -205,7 +207,9 @@ async def test_should_return_model_not_found_when_models_response_is_empty(self, assert result == ModelNotFoundError(name=DEFAULT_MODEL_ID) @pytest.mark.asyncio - async def test_should_return_model_not_found_when_model_is_missing_in_models_response(self, repository: ProviderCapabilitiesRepository, provider_client: Mock): + async def test_should_return_model_not_found_when_model_is_missing_in_models_response( + self, repository: ProviderCapabilitiesRepository, provider_client: Mock + ): provider_client.forward_request.return_value = ProviderOriginalResponse(data=AlbertModelsResponseFactory(data=[AlbertModelResponseFactory()])) result = await repository._get_max_context_length(adapter=models_adapter()) @@ -213,7 +217,9 @@ async def test_should_return_model_not_found_when_model_is_missing_in_models_res assert result == ModelNotFoundError(name=DEFAULT_MODEL_ID) @pytest.mark.asyncio - async def test_should_return_provider_not_reachable_when_getting_max_context_fails(self, repository: ProviderCapabilitiesRepository, provider_client: Mock): + async def test_should_return_provider_not_reachable_when_getting_max_context_fails( + self, repository: ProviderCapabilitiesRepository, provider_client: Mock + ): provider_client.forward_request.return_value = StatusCodeModelError(status_code=500, detail="boom") result = await repository._get_max_context_length(adapter=models_adapter()) @@ -236,7 +242,9 @@ async def test_should_get_vector_size(self, repository: ProviderCapabilitiesRepo assert formatted_request.body["model"] == DEFAULT_MODEL_ID @pytest.mark.asyncio - async def test_should_return_provider_not_reachable_when_getting_vector_size_fails(self, repository: ProviderCapabilitiesRepository, provider_client: Mock): + async def test_should_return_provider_not_reachable_when_getting_vector_size_fails( + self, repository: ProviderCapabilitiesRepository, provider_client: Mock + ): provider_client.forward_request.return_value = StatusCodeModelError(status_code=500, detail="boom") result = await repository._get_vector_size(adapter=embeddings_adapter()) diff --git a/api/tests/unit/use_case/admin/providers/test_createproviderusecase.py b/api/tests/unit/use_case/admin/providers/test_createproviderusecase.py index f75775804..af97bccd4 100644 --- a/api/tests/unit/use_case/admin/providers/test_createproviderusecase.py +++ b/api/tests/unit/use_case/admin/providers/test_createproviderusecase.py @@ -20,6 +20,7 @@ def router_repository(): def provider_repository(): return AsyncMock() + @pytest.fixture def provider_capabilities_repository(): return AsyncMock() @@ -353,7 +354,9 @@ async def test_should_return_provider_not_reachable_error_when_gateway_fails( # Arrange router_repository.get_router_by_id.return_value = sample_router - provider_capabilities_repository.get_provider_capabilities.return_value = ProviderNotReachableError(model_name="my-model", status_code=500, detail="error_detail") + provider_capabilities_repository.get_provider_capabilities.return_value = ProviderNotReachableError( + model_name="my-model", status_code=500, detail="error_detail" + ) # Act result = await use_case.execute(default_command) diff --git a/api/tests/unit/use_case/admin/test_bootstrapmodelsusecase.py b/api/tests/unit/use_case/admin/test_bootstrapmodelsusecase.py index 6cf2ce002..880c3d4e2 100644 --- a/api/tests/unit/use_case/admin/test_bootstrapmodelsusecase.py +++ b/api/tests/unit/use_case/admin/test_bootstrapmodelsusecase.py @@ -35,7 +35,11 @@ def provider_capabilities_repository(): @pytest.fixture def use_case(router_repository, provider_repository, provider_capabilities_repository): - return BootstrapModelsUseCase(router_repository=router_repository, provider_repository=provider_repository, provider_capabilities_repository=provider_capabilities_repository) + return BootstrapModelsUseCase( + router_repository=router_repository, + provider_repository=provider_repository, + provider_capabilities_repository=provider_capabilities_repository, + ) class TestBootstrapModelsUseCase: @@ -55,7 +59,9 @@ async def test_skips_when_routers_already_exist(self, use_case, router_repositor provider_repository.create_provider.assert_not_awaited() @pytest.mark.asyncio - async def test_successfully_creates_router_with_single_provider(self, use_case, router_repository, provider_repository, provider_capabilities_repository): + async def test_successfully_creates_router_with_single_provider( + self, use_case, router_repository, provider_repository, provider_capabilities_repository + ): # Arrange model_provider = ModelProviderConfigurationFactory() model_configuration = ModelConfigurationFactory(providers=[model_provider]) @@ -241,13 +247,17 @@ async def test_returns_provider_already_exists_error_when_duplicate_within_route provider_repository.create_provider.assert_not_awaited() @pytest.mark.asyncio - async def test_returns_provider_not_reachable_error_and_rolls_back(self, use_case, router_repository, provider_repository, provider_capabilities_repository): + async def test_returns_provider_not_reachable_error_and_rolls_back( + self, use_case, router_repository, provider_repository, provider_capabilities_repository + ): # Arrange model_configuration = ModelConfigurationFactory() router = RouterFactory(id=1, name=model_configuration.name, type=RouterType.TEXT_GENERATION) router_repository.get_all_routers.return_value = [] router_repository.create_router.return_value = router - provider_capabilities_repository.get_provider_capabilities.return_value = ProviderNotReachableError(model_name="my-model", status_code=500, detail="error_detail") + provider_capabilities_repository.get_provider_capabilities.return_value = ProviderNotReachableError( + model_name="my-model", status_code=500, detail="error_detail" + ) # Act result = await use_case.execute(routers_to_create=[model_configuration], bootstrap_admin_user_id=BOOTSTRAP_ADMIN_USER_ID) @@ -258,7 +268,9 @@ async def test_returns_provider_not_reachable_error_and_rolls_back(self, use_cas router_repository.delete_all_routers.assert_awaited_once() @pytest.mark.asyncio - async def test_returns_model_not_found_error_and_rolls_back(self, use_case, router_repository, provider_repository, provider_capabilities_repository): + async def test_returns_model_not_found_error_and_rolls_back( + self, use_case, router_repository, provider_repository, provider_capabilities_repository + ): # Arrange model_configuration = ModelConfigurationFactory() router = RouterFactory(id=1, name=model_configuration.name, type=RouterType.TEXT_GENERATION) @@ -309,7 +321,9 @@ async def test_returns_inconsistent_max_context_length_error_and_rolls_back( router_repository.delete_all_routers.assert_awaited_once() @pytest.mark.asyncio - async def test_returns_inconsistent_vector_size_error_and_rolls_back(self, use_case, router_repository, provider_repository, provider_capabilities_repository): + async def test_returns_inconsistent_vector_size_error_and_rolls_back( + self, use_case, router_repository, provider_repository, provider_capabilities_repository + ): # Arrange model_configuration = ModelConfigurationFactory( type=RouterType.TEXT_EMBEDDINGS_INFERENCE, @@ -345,7 +359,9 @@ async def test_returns_inconsistent_vector_size_error_and_rolls_back(self, use_c router_repository.delete_all_routers.assert_awaited_once() @pytest.mark.asyncio - async def test_returns_success_with_no_routers_to_create(self, use_case, router_repository, provider_repository, provider_capabilities_repository): + async def test_returns_success_with_no_routers_to_create( + self, use_case, router_repository, provider_repository, provider_capabilities_repository + ): # Arrange router_repository.get_all_routers.return_value = [] diff --git a/api/use_cases/admin/providers/_createproviderusecase.py b/api/use_cases/admin/providers/_createproviderusecase.py index ff1d66aec..79d1f4ea6 100644 --- a/api/use_cases/admin/providers/_createproviderusecase.py +++ b/api/use_cases/admin/providers/_createproviderusecase.py @@ -43,7 +43,12 @@ class CreateProviderUseCaseSuccess: class CreateProviderUseCase: - def __init__(self, router_repository: RouterRepository, provider_repository: ProviderRepository, provider_capabilities_repository: ProviderCapabilitiesRepository): + def __init__( + self, + router_repository: RouterRepository, + provider_repository: ProviderRepository, + provider_capabilities_repository: ProviderCapabilitiesRepository, + ): self.router_repository = router_repository self.provider_repository = provider_repository self.provider_capabilities_repository = provider_capabilities_repository diff --git a/api/use_cases/models/_bootstrapmodelsusecase.py b/api/use_cases/models/_bootstrapmodelsusecase.py index 53f73940f..38d37d07c 100644 --- a/api/use_cases/models/_bootstrapmodelsusecase.py +++ b/api/use_cases/models/_bootstrapmodelsusecase.py @@ -35,7 +35,12 @@ class BootstrapModelsUseCaseSkipped: class BootstrapModelsUseCase: - def __init__(self, router_repository: RouterRepository, provider_repository: ProviderRepository, provider_capabilities_repository: ProviderCapabilitiesRepository): + def __init__( + self, + router_repository: RouterRepository, + provider_repository: ProviderRepository, + provider_capabilities_repository: ProviderCapabilitiesRepository, + ): self.router_repository = router_repository self.provider_repository = provider_repository self.provider_capabilities_repository = provider_capabilities_repository From c4a40516351798e736e236b9715d1c1350533133 Mon Sep 17 00:00:00 2001 From: bakr-a Date: Thu, 23 Jul 2026 14:25:01 +0000 Subject: [PATCH 15/17] Update unit coverage badge --- .github/badges/coverage.json | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/.github/badges/coverage.json b/.github/badges/coverage.json index 68cd04f7e..dea2fde8e 100644 --- a/.github/badges/coverage.json +++ b/.github/badges/coverage.json @@ -1 +1 @@ -{"schemaVersion": 1, "label": "coverage", "message": "59.82%", "color": "red"} \ No newline at end of file +{"schemaVersion": 1, "label": "coverage", "message": "59.78%", "color": "red"} \ No newline at end of file From b8a8dd9f3d74c93720e5c160590d853c5d324b86 Mon Sep 17 00:00:00 2001 From: Bakr Annour Date: Thu, 23 Jul 2026 16:40:30 +0200 Subject: [PATCH 16/17] docs(adr): replace ProviderGateway with ProviderCapabilitiesRepository Co-Authored-By: Claude Fable 5 --- adr/2026-05-28-refactoring-model-forwarding.md | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/adr/2026-05-28-refactoring-model-forwarding.md b/adr/2026-05-28-refactoring-model-forwarding.md index 16274e3da..eeb4b96bf 100644 --- a/adr/2026-05-28-refactoring-model-forwarding.md +++ b/adr/2026-05-28-refactoring-model-forwarding.md @@ -109,6 +109,7 @@ The new architecture follows the principles of clean architecture: the model for * **Domain contracts for forwarding:** provider client, provider load balancer, router rate limiter, model tokenizer and environmental impact computer are exposed as domain abstractions. The use case depends on these contracts, not on Redis, HTTP, Ecologit or Tiktoken directly. * **HTTP client simplified:** `HttpProviderClient` only sends an already formatted request to the selected provider and returns the raw provider response or a model error. It no longer owns endpoint selection, usage computation, metrics or rate limiting. * **Endpoint adapters extracted:** provider-specific adapters convert OpenGate requests and responses to each provider format. `build_adapter` selects the right adapter from the source endpoint and provider type, while common usage computation stays in the base adapter. +* **Provider gateway removed:** the `ProviderGateway` contract and its `ModelProviderGateway` infrastructure implementation are deleted. Provider capability fetching (max context length, vector size) now lives in `ProviderCapabilitiesRepository`, a domain service that composes the `ProviderClient` and `ProviderAdapterBuilder` contracts. * **Redis responsibilities isolated:** Redis implementations handle provider load balancing, provider metrics and router rate limits behind dedicated contracts. The use case decides when those operations happen. * **Usage and impacts made explicit:** prompt tokens are computed before rate limiting, response usage is computed after provider response formatting, and environmental impacts are delegated to the Ecologit implementation through a domain contract. * **FastAPI endpoint thinned:** the HTTP endpoint builds the command, calls the use case and maps domain errors to HTTP exceptions. It no longer contains forwarding logic. @@ -134,7 +135,7 @@ subgraph DL[**Domain layer**] provider_adapter_builder[ProviderAdapterBuilder] provider_adapter[ProviderAdapter] provider_repository[ProviderRepository] - provider_gateway[ProviderGateway] + provider_capabilities_repository[ProviderCapabilitiesRepository] provider_load_balancer[ProviderLoadBalancer] provider_client[ProviderClient] provider_metrics_logger[ProviderMetricsLogger] @@ -187,7 +188,6 @@ use_case --> router_rate_limiter use_case --> provider_repository use_case --> provider_adapter_builder use_case --> provider_adapter -use_case --> provider_gateway use_case --> provider_load_balancer use_case --> provider_client use_case --> provider_metrics_logger @@ -197,6 +197,8 @@ use_case --> user_with_role_query use_case --> usage_computer usage_computer --> model_environmental_impacts_computer usage_computer --> model_tokenizer +provider_capabilities_repository --> provider_client +provider_capabilities_repository --> provider_adapter_builder From c599dcc04181600fa5fef1b01f78255a2e2c0540 Mon Sep 17 00:00:00 2001 From: Benjamin PILIA Date: Tue, 28 Jul 2026 10:34:54 +0200 Subject: [PATCH 17/17] review: move files --- ...2026-05-28-refactoring-model-forwarding.md | 7 +- api/dependencies.py | 12 +-- api/domain/provider/__init__.py | 2 - .../admin/provider/test_create_provider.py | 7 +- .../test_providercapabilitiesprobe.py} | 16 ++-- .../providers/test_createproviderusecase.py | 48 ++++++------ .../admin/test_bootstrapmodelsusecase.py | 54 ++++++------- .../services}/__init__.py | 0 .../test_providercapabilitiesprobe.py} | 76 +++++++++---------- .../admin/providers/_createproviderusecase.py | 9 ++- .../models/_bootstrapmodelsusecase.py | 9 ++- api/use_cases/services/__init__.py | 3 + .../services/_providercapabilitiesprobe.py} | 8 +- api/utils/lifespan.py | 6 +- 14 files changed, 124 insertions(+), 133 deletions(-) rename api/tests/integration/{http/test_providercapabilitiesrepository.py => use_case/test_providercapabilitiesprobe.py} (85%) rename api/tests/unit/{domain/provider => use_case/services}/__init__.py (100%) rename api/tests/unit/{domain/provider/test_providercapabilitiesrepository.py => use_case/services/test_providercapabilitiesprobe.py} (74%) create mode 100644 api/use_cases/services/__init__.py rename api/{domain/provider/_providercapabilitiesrepository.py => use_cases/services/_providercapabilitiesprobe.py} (93%) diff --git a/adr/2026-05-28-refactoring-model-forwarding.md b/adr/2026-05-28-refactoring-model-forwarding.md index eeb4b96bf..4b581f0d6 100644 --- a/adr/2026-05-28-refactoring-model-forwarding.md +++ b/adr/2026-05-28-refactoring-model-forwarding.md @@ -109,7 +109,7 @@ The new architecture follows the principles of clean architecture: the model for * **Domain contracts for forwarding:** provider client, provider load balancer, router rate limiter, model tokenizer and environmental impact computer are exposed as domain abstractions. The use case depends on these contracts, not on Redis, HTTP, Ecologit or Tiktoken directly. * **HTTP client simplified:** `HttpProviderClient` only sends an already formatted request to the selected provider and returns the raw provider response or a model error. It no longer owns endpoint selection, usage computation, metrics or rate limiting. * **Endpoint adapters extracted:** provider-specific adapters convert OpenGate requests and responses to each provider format. `build_adapter` selects the right adapter from the source endpoint and provider type, while common usage computation stays in the base adapter. -* **Provider gateway removed:** the `ProviderGateway` contract and its `ModelProviderGateway` infrastructure implementation are deleted. Provider capability fetching (max context length, vector size) now lives in `ProviderCapabilitiesRepository`, a domain service that composes the `ProviderClient` and `ProviderAdapterBuilder` contracts. +* **Provider gateway replaced by a capabilities probe:** the `ProviderGateway` contract and its `ModelProviderGateway` implementation are replaced by a single `ProviderCapabilitiesProbe` class in the use case layer. Probing a provider for its capabilities (max context length, vector size) is an acquisition concern, not a repository. The class holds no technology of its own: it composes the `ProviderClient` and `ProviderAdapterBuilder` domain contracts and sequences two calls, which makes it an application service shared by `CreateProviderUseCase` and `BootstrapModelsUseCase` rather than an adapter. No abstraction is declared for it, since it is not a boundary the domain needs to invert. It does not appear in the diagram below: the model forwarding use case does not probe capabilities. * **Redis responsibilities isolated:** Redis implementations handle provider load balancing, provider metrics and router rate limits behind dedicated contracts. The use case decides when those operations happen. * **Usage and impacts made explicit:** prompt tokens are computed before rate limiting, response usage is computed after provider response formatting, and environmental impacts are delegated to the Ecologit implementation through a domain contract. * **FastAPI endpoint thinned:** the HTTP endpoint builds the command, calls the use case and maps domain errors to HTTP exceptions. It no longer contains forwarding logic. @@ -135,7 +135,6 @@ subgraph DL[**Domain layer**] provider_adapter_builder[ProviderAdapterBuilder] provider_adapter[ProviderAdapter] provider_repository[ProviderRepository] - provider_capabilities_repository[ProviderCapabilitiesRepository] provider_load_balancer[ProviderLoadBalancer] provider_client[ProviderClient] provider_metrics_logger[ProviderMetricsLogger] @@ -197,10 +196,6 @@ use_case --> user_with_role_query use_case --> usage_computer usage_computer --> model_environmental_impacts_computer usage_computer --> model_tokenizer -provider_capabilities_repository --> provider_client -provider_capabilities_repository --> provider_adapter_builder - - provider_adapter_builder --> http_provider_adapter_builder provider_client --> http_provider_client diff --git a/api/dependencies.py b/api/dependencies.py index 15713a6d1..7d7e57321 100644 --- a/api/dependencies.py +++ b/api/dependencies.py @@ -11,7 +11,6 @@ from api.domain.model import ModelEnvironmentalImpactsComputer, ModelTokenizer from api.domain.provider import ( ProviderAdapterBuilder, - ProviderCapabilitiesRepository, ProviderClient, ProviderLoadBalancer, ProviderMetricsLogger, @@ -55,6 +54,7 @@ from api.use_cases.health import GetHealthModelsUseCase from api.use_cases.models import GetModelsUseCase, GetModelUseCase from api.use_cases.reranks import CreateRerankUseCase +from api.use_cases.services import ProviderCapabilitiesProbe from api.utils.configuration import configuration from api.utils.context import global_context @@ -162,11 +162,11 @@ def _provider_repository(session: AsyncSession) -> ProviderRepository: return PostgresProviderRepository(postgres_session=session) -def _provider_capabilities_repository( +def _provider_capabilities_probe( provider_client: ProviderClient = Depends(_provider_client), provider_adapter_builder: ProviderAdapterBuilder = Depends(_provider_adapter_builder), -) -> ProviderCapabilitiesRepository: - return ProviderCapabilitiesRepository(provider_client=provider_client, provider_adapter_builder=provider_adapter_builder) +) -> ProviderCapabilitiesProbe: + return ProviderCapabilitiesProbe(provider_client=provider_client, provider_adapter_builder=provider_adapter_builder) # health use cases @@ -361,12 +361,12 @@ def update_router_use_case_factory(postgres_session: AsyncSession = Depends(get_ # provider use cases def create_provider_use_case_factory( postgres_session: AsyncSession = Depends(get_postgres_session), - provider_capabilities_repository: ProviderCapabilitiesRepository = Depends(_provider_capabilities_repository), + provider_capabilities_probe: ProviderCapabilitiesProbe = Depends(_provider_capabilities_probe), ) -> CreateProviderUseCase: return CreateProviderUseCase( router_repository=_router_repository(postgres_session), provider_repository=_provider_repository(postgres_session), - provider_capabilities_repository=provider_capabilities_repository, + provider_capabilities_probe=provider_capabilities_probe, ) diff --git a/api/domain/provider/__init__.py b/api/domain/provider/__init__.py index fe90bcb67..8e80d707c 100644 --- a/api/domain/provider/__init__.py +++ b/api/domain/provider/__init__.py @@ -1,6 +1,5 @@ from api.domain.provider._provideradapter import ProviderAdapter from api.domain.provider._provideradapterbuilder import ProviderAdapterBuilder -from api.domain.provider._providercapabilitiesrepository import ProviderCapabilitiesRepository from api.domain.provider._providerclient import ProviderClient, ProviderClientResponse from api.domain.provider._providerloadbalancer import ProviderLoadBalancer from api.domain.provider._providermetricslogger import ProviderMetricsLogger @@ -14,5 +13,4 @@ "ProviderLoadBalancer", "ProviderMetricsLogger", "ProviderRepository", - "ProviderCapabilitiesRepository", ] diff --git a/api/tests/integration/endpoints/admin/provider/test_create_provider.py b/api/tests/integration/endpoints/admin/provider/test_create_provider.py index 0bc640583..295e1e9ee 100644 --- a/api/tests/integration/endpoints/admin/provider/test_create_provider.py +++ b/api/tests/integration/endpoints/admin/provider/test_create_provider.py @@ -7,7 +7,7 @@ from api.dependencies import create_provider_use_case_factory from api.domain.model.entities import ModelType as RouterType -from api.domain.model.errors import InconsistentModelMaxContextLengthError, InconsistentModelVectorSizeError +from api.domain.model.errors import InconsistentModelMaxContextLengthError, InconsistentModelVectorSizeError, ModelNotFoundError from api.domain.provider.entities import ProviderType from api.domain.provider.errors import InvalidProviderTypeError, ProviderAlreadyExistsError, ProviderNotReachableError from api.domain.router.errors import RouterNotFoundError @@ -84,6 +84,11 @@ async def test_happy_path(self, client: AsyncClient, db_session): 424, "Model provider my-model not reachable (500): error_detail", ), + ( + ModelNotFoundError(name="my-model"), + 404, + "Model my-model not found.", + ), ( ProviderAlreadyExistsError(model_name="my-model", url=DEFAULT_PROVIDER_URL, router_id=1), 409, diff --git a/api/tests/integration/http/test_providercapabilitiesrepository.py b/api/tests/integration/use_case/test_providercapabilitiesprobe.py similarity index 85% rename from api/tests/integration/http/test_providercapabilitiesrepository.py rename to api/tests/integration/use_case/test_providercapabilitiesprobe.py index 0597dc295..f7be808ac 100644 --- a/api/tests/integration/http/test_providercapabilitiesrepository.py +++ b/api/tests/integration/use_case/test_providercapabilitiesprobe.py @@ -5,11 +5,11 @@ import respx from api.domain.model.entities import ModelType as RouterType -from api.domain.provider import ProviderCapabilitiesRepository from api.domain.provider.entities import ProviderCapabilities, ProviderType from api.infrastructure.http import HttpProviderAdapterBuilder, HttpProviderClient from api.tests.integration.factories.tei import TeiEmbeddingsResponseFactory, TeiModelsResponseFactory from api.tests.integration.factories.vllm import VllmModelsResponseFactory +from api.use_cases.services import ProviderCapabilitiesProbe DEFAULT_PROVIDER_URL = "http://my-test-provider/" DEFAULT_MODEL_ID = "test/my-model" @@ -31,14 +31,14 @@ def _mock_embeddings_response(respx_mock, body: dict, status_code: int) -> None: @pytest.fixture -def repository() -> ProviderCapabilitiesRepository: - return ProviderCapabilitiesRepository(provider_client=HttpProviderClient(), provider_adapter_builder=HttpProviderAdapterBuilder()) +def probe() -> ProviderCapabilitiesProbe: + return ProviderCapabilitiesProbe(provider_client=HttpProviderClient(), provider_adapter_builder=HttpProviderAdapterBuilder()) @pytest.mark.asyncio(loop_scope="session") -class TestProviderCapabilitiesRepository: +class TestProviderCapabilitiesProbe: @respx.mock - async def test_get_capabilities_of_non_embeddings_providers(self, repository: ProviderCapabilitiesRepository): + async def test_get_capabilities_of_non_embeddings_providers(self, probe: ProviderCapabilitiesProbe): _mock_models_response( respx_mock=respx, provider_type=ProviderType.VLLM, @@ -46,7 +46,7 @@ async def test_get_capabilities_of_non_embeddings_providers(self, repository: Pr status_code=VllmModelsResponseFactory._status_code, ) - result = await repository.get_provider_capabilities( + result = await probe.get_capabilities( router_type=RouterType.TEXT_GENERATION, provider_type=ProviderType.VLLM, url=DEFAULT_PROVIDER_URL, @@ -60,7 +60,7 @@ async def test_get_capabilities_of_non_embeddings_providers(self, repository: Pr @respx.mock async def test_get_capabilities_of_embeddings_providers( self, - repository: ProviderCapabilitiesRepository, + probe: ProviderCapabilitiesProbe, ): _mock_models_response( respx_mock=respx, @@ -74,7 +74,7 @@ async def test_get_capabilities_of_embeddings_providers( status_code=TeiEmbeddingsResponseFactory._status_code, ) - result = await repository.get_provider_capabilities( + result = await probe.get_capabilities( router_type=RouterType.TEXT_EMBEDDINGS_INFERENCE, provider_type=ProviderType.TEI, url=DEFAULT_PROVIDER_URL, diff --git a/api/tests/unit/use_case/admin/providers/test_createproviderusecase.py b/api/tests/unit/use_case/admin/providers/test_createproviderusecase.py index af97bccd4..d8e2e6388 100644 --- a/api/tests/unit/use_case/admin/providers/test_createproviderusecase.py +++ b/api/tests/unit/use_case/admin/providers/test_createproviderusecase.py @@ -22,16 +22,16 @@ def provider_repository(): @pytest.fixture -def provider_capabilities_repository(): +def provider_capabilities_probe(): return AsyncMock() @pytest.fixture -def use_case(router_repository, provider_repository, provider_capabilities_repository): +def use_case(router_repository, provider_repository, provider_capabilities_probe): return CreateProviderUseCase( router_repository=router_repository, provider_repository=provider_repository, - provider_capabilities_repository=provider_capabilities_repository, + provider_capabilities_probe=provider_capabilities_probe, ) @@ -168,7 +168,7 @@ async def test_should_create_provider_when_router_has_a_different_provider( use_case, router_repository, provider_repository, - provider_capabilities_repository, + provider_capabilities_probe, sample_router_with_providers, sample_provider, default_command, @@ -176,7 +176,7 @@ async def test_should_create_provider_when_router_has_a_different_provider( # Arrange router_repository.get_router_by_id.return_value = sample_router_with_providers - provider_capabilities_repository.get_provider_capabilities.return_value = ProviderCapabilities(max_context_length=4096, vector_size=None) + provider_capabilities_probe.get_capabilities.return_value = ProviderCapabilities(max_context_length=4096, vector_size=None) provider_repository.create_provider.return_value = sample_provider # Act @@ -209,7 +209,7 @@ async def test_should_create_embedding_provider_when_vector_size_matches( use_case, router_repository, provider_repository, - provider_capabilities_repository, + provider_capabilities_probe, sample_embedding_router_with_providers, sample_provider, default_command, @@ -217,7 +217,7 @@ async def test_should_create_embedding_provider_when_vector_size_matches( # Arrange router_repository.get_router_by_id.return_value = sample_embedding_router_with_providers - provider_capabilities_repository.get_provider_capabilities.return_value = ProviderCapabilities(max_context_length=512, vector_size=768) + provider_capabilities_probe.get_capabilities.return_value = ProviderCapabilities(max_context_length=512, vector_size=768) provider_repository.create_provider.return_value = sample_provider # Act @@ -246,7 +246,7 @@ async def test_should_create_embedding_provider_when_vector_size_matches( @pytest.mark.asyncio async def test_should_return_router_not_found_error_when_router_does_not_exist( - self, use_case, router_repository, provider_repository, provider_capabilities_repository, default_command + self, use_case, router_repository, provider_repository, provider_capabilities_probe, default_command ): # Arrange @@ -258,7 +258,7 @@ async def test_should_return_router_not_found_error_when_router_does_not_exist( # Assert assert isinstance(result, RouterNotFoundError) assert result.id == 1 - provider_capabilities_repository.get_provider_capabilities.assert_not_called() + provider_capabilities_probe.get_capabilities.assert_not_called() provider_repository.create_provider.assert_not_called() @pytest.mark.asyncio @@ -272,7 +272,7 @@ async def test_should_create_provider_when_provider_type_is_compatible( use_case, router_repository, provider_repository, - provider_capabilities_repository, + provider_capabilities_probe, default_command, router_type, provider_type, @@ -281,7 +281,7 @@ async def test_should_create_provider_when_provider_type_is_compatible( capabilities = capabilities_for(router_type) provider = ProviderFactory(id=1, router_id=1, user_id=1, type=provider_type, url="https://example.com/", model_name="my-model") router_repository.get_router_by_id.return_value = RouterFactory(id=1, name="test-router", type=router_type, providers=0) - provider_capabilities_repository.get_provider_capabilities.return_value = capabilities + provider_capabilities_probe.get_capabilities.return_value = capabilities provider_repository.create_provider.return_value = provider command = with_provider_type(default_command, provider_type) @@ -291,7 +291,7 @@ async def test_should_create_provider_when_provider_type_is_compatible( # Assert assert isinstance(result, CreateProviderUseCaseSuccess) assert result.provider == provider - provider_capabilities_repository.get_provider_capabilities.assert_called_once_with( + provider_capabilities_probe.get_capabilities.assert_called_once_with( router_type=router_type, provider_type=provider_type, url="https://example.com/", @@ -328,7 +328,7 @@ async def test_should_return_invalid_provider_type_error_when_provider_type_is_n use_case, router_repository, provider_repository, - provider_capabilities_repository, + provider_capabilities_probe, default_command, router_type, provider_type, @@ -344,17 +344,17 @@ async def test_should_return_invalid_provider_type_error_when_provider_type_is_n assert isinstance(result, InvalidProviderTypeError) assert result.provider_type == provider_type.value assert result.router_type == router_type.value - provider_capabilities_repository.get_provider_capabilities.assert_not_called() + provider_capabilities_probe.get_capabilities.assert_not_called() provider_repository.create_provider.assert_not_called() @pytest.mark.asyncio async def test_should_return_provider_not_reachable_error_when_gateway_fails( - self, use_case, router_repository, provider_repository, provider_capabilities_repository, sample_router, default_command + self, use_case, router_repository, provider_repository, provider_capabilities_probe, sample_router, default_command ): # Arrange router_repository.get_router_by_id.return_value = sample_router - provider_capabilities_repository.get_provider_capabilities.return_value = ProviderNotReachableError( + provider_capabilities_probe.get_capabilities.return_value = ProviderNotReachableError( model_name="my-model", status_code=500, detail="error_detail" ) @@ -370,12 +370,12 @@ async def test_should_return_provider_not_reachable_error_when_gateway_fails( @pytest.mark.asyncio async def test_should_return_model_not_found_error_when_model_is_missing( - self, use_case, router_repository, provider_repository, provider_capabilities_repository, sample_router, default_command + self, use_case, router_repository, provider_repository, provider_capabilities_probe, sample_router, default_command ): # Arrange router_repository.get_router_by_id.return_value = sample_router - provider_capabilities_repository.get_provider_capabilities.return_value = ModelNotFoundError(name="my-model") + provider_capabilities_probe.get_capabilities.return_value = ModelNotFoundError(name="my-model") # Act result = await use_case.execute(default_command) @@ -387,12 +387,12 @@ async def test_should_return_model_not_found_error_when_model_is_missing( @pytest.mark.asyncio async def test_should_return_inconsistent_max_context_length_error_when_mismatch( - self, use_case, router_repository, provider_repository, provider_capabilities_repository, sample_router_with_providers, default_command + self, use_case, router_repository, provider_repository, provider_capabilities_probe, sample_router_with_providers, default_command ): # Arrange router_repository.get_router_by_id.return_value = sample_router_with_providers - provider_capabilities_repository.get_provider_capabilities.return_value = ProviderCapabilities(max_context_length=2048, vector_size=None) + provider_capabilities_probe.get_capabilities.return_value = ProviderCapabilities(max_context_length=2048, vector_size=None) # Act result = await use_case.execute(default_command) @@ -409,14 +409,14 @@ async def test_should_return_inconsistent_vector_size_error_when_mismatch( use_case, router_repository, provider_repository, - provider_capabilities_repository, + provider_capabilities_probe, sample_embedding_router_with_providers, default_command, ): # Arrange router_repository.get_router_by_id.return_value = sample_embedding_router_with_providers - provider_capabilities_repository.get_provider_capabilities.return_value = ProviderCapabilities(max_context_length=512, vector_size=384) + provider_capabilities_probe.get_capabilities.return_value = ProviderCapabilities(max_context_length=512, vector_size=384) # Act result = await use_case.execute(with_provider_type(default_command, ProviderType.TEI)) @@ -429,12 +429,12 @@ async def test_should_return_inconsistent_vector_size_error_when_mismatch( @pytest.mark.asyncio async def test_should_return_provider_already_exists_error( - self, use_case, router_repository, provider_repository, provider_capabilities_repository, sample_router, default_command + self, use_case, router_repository, provider_repository, provider_capabilities_probe, sample_router, default_command ): # Arrange router_repository.get_router_by_id.return_value = sample_router - provider_capabilities_repository.get_provider_capabilities.return_value = ProviderCapabilities(max_context_length=4096, vector_size=None) + provider_capabilities_probe.get_capabilities.return_value = ProviderCapabilities(max_context_length=4096, vector_size=None) provider_repository.create_provider.return_value = ProviderAlreadyExistsError(model_name="my-model", url="https://example.com/", router_id=1) # Act diff --git a/api/tests/unit/use_case/admin/test_bootstrapmodelsusecase.py b/api/tests/unit/use_case/admin/test_bootstrapmodelsusecase.py index 880c3d4e2..37fb9abd8 100644 --- a/api/tests/unit/use_case/admin/test_bootstrapmodelsusecase.py +++ b/api/tests/unit/use_case/admin/test_bootstrapmodelsusecase.py @@ -29,22 +29,22 @@ def provider_repository(): @pytest.fixture -def provider_capabilities_repository(): +def provider_capabilities_probe(): return AsyncMock() @pytest.fixture -def use_case(router_repository, provider_repository, provider_capabilities_repository): +def use_case(router_repository, provider_repository, provider_capabilities_probe): return BootstrapModelsUseCase( router_repository=router_repository, provider_repository=provider_repository, - provider_capabilities_repository=provider_capabilities_repository, + provider_capabilities_probe=provider_capabilities_probe, ) class TestBootstrapModelsUseCase: @pytest.mark.asyncio - async def test_skips_when_routers_already_exist(self, use_case, router_repository, provider_repository, provider_capabilities_repository): + async def test_skips_when_routers_already_exist(self, use_case, router_repository, provider_repository, provider_capabilities_probe): # Arrange existing_routers = [RouterFactory(id=1), RouterFactory(id=2)] router_repository.get_all_routers.return_value = existing_routers @@ -55,12 +55,12 @@ async def test_skips_when_routers_already_exist(self, use_case, router_repositor # Assert assert result == BootstrapModelsUseCaseSkipped(number_of_routers=2) router_repository.create_router.assert_not_awaited() - provider_capabilities_repository.get_provider_capabilities.assert_not_awaited() + provider_capabilities_probe.get_capabilities.assert_not_awaited() provider_repository.create_provider.assert_not_awaited() @pytest.mark.asyncio async def test_successfully_creates_router_with_single_provider( - self, use_case, router_repository, provider_repository, provider_capabilities_repository + self, use_case, router_repository, provider_repository, provider_capabilities_probe ): # Arrange model_provider = ModelProviderConfigurationFactory() @@ -70,7 +70,7 @@ async def test_successfully_creates_router_with_single_provider( router_repository.get_all_routers.return_value = [] router_repository.create_router.return_value = router - provider_capabilities_repository.get_provider_capabilities.return_value = ProviderCapabilities(max_context_length=4096, vector_size=None) + provider_capabilities_probe.get_capabilities.return_value = ProviderCapabilities(max_context_length=4096, vector_size=None) provider_repository.create_provider.return_value = provider # Act @@ -88,7 +88,7 @@ async def test_successfully_creates_router_with_single_provider( user_id=BOOTSTRAP_ADMIN_USER_ID, aliases=model_configuration.aliases, ) - provider_capabilities_repository.get_provider_capabilities.assert_awaited_once_with( + provider_capabilities_probe.get_capabilities.assert_awaited_once_with( router_type=router.type, provider_type=model_provider.type, url=model_provider.url, @@ -117,7 +117,7 @@ async def test_successfully_creates_router_with_single_provider( @pytest.mark.asyncio async def test_successfully_creates_multiple_routers_with_multiple_providers( - self, use_case, router_repository, provider_repository, provider_capabilities_repository + self, use_case, router_repository, provider_repository, provider_capabilities_probe ): # Arrange first_model = ModelConfigurationFactory( @@ -139,7 +139,7 @@ async def test_successfully_creates_multiple_routers_with_multiple_providers( router_repository.get_all_routers.return_value = [] router_repository.create_router.side_effect = [first_router, second_router] - provider_capabilities_repository.get_provider_capabilities.side_effect = [ + provider_capabilities_probe.get_capabilities.side_effect = [ ProviderCapabilities(max_context_length=4096, vector_size=None), ProviderCapabilities(max_context_length=4096, vector_size=None), ProviderCapabilities(max_context_length=512, vector_size=768), @@ -159,13 +159,13 @@ async def test_successfully_creates_multiple_routers_with_multiple_providers( # Assert assert result == BootstrapModelsUseCaseSuccess(number_of_routers=2) assert router_repository.create_router.await_count == 2 - assert provider_capabilities_repository.get_provider_capabilities.await_count == 3 + assert provider_capabilities_probe.get_capabilities.await_count == 3 assert provider_repository.create_provider.await_count == 3 router_repository.delete_all_routers.assert_not_awaited() @pytest.mark.asyncio async def test_returns_router_name_already_exists_error_when_duplicate_name( - self, use_case, router_repository, provider_repository, provider_capabilities_repository + self, use_case, router_repository, provider_repository, provider_capabilities_probe ): # Arrange routers_to_create = [ @@ -183,7 +183,7 @@ async def test_returns_router_name_already_exists_error_when_duplicate_name( # Assert assert result == RouterNameAlreadyExistsError(name="duplicate") router_repository.create_router.assert_not_awaited() - provider_capabilities_repository.get_provider_capabilities.assert_not_awaited() + provider_capabilities_probe.get_capabilities.assert_not_awaited() provider_repository.create_provider.assert_not_awaited() @pytest.mark.asyncio @@ -224,7 +224,7 @@ async def test_returns_provider_already_exists_error_when_duplicate_within_route use_case, router_repository, provider_repository, - provider_capabilities_repository, + provider_capabilities_probe, ): # Arrange model_configuration = ModelConfigurationFactory( @@ -243,19 +243,19 @@ async def test_returns_provider_already_exists_error_when_duplicate_within_route assert result.model_name == "model-a" assert result.url == "https://provider.com/" router_repository.create_router.assert_not_awaited() - provider_capabilities_repository.get_provider_capabilities.assert_not_awaited() + provider_capabilities_probe.get_capabilities.assert_not_awaited() provider_repository.create_provider.assert_not_awaited() @pytest.mark.asyncio async def test_returns_provider_not_reachable_error_and_rolls_back( - self, use_case, router_repository, provider_repository, provider_capabilities_repository + self, use_case, router_repository, provider_repository, provider_capabilities_probe ): # Arrange model_configuration = ModelConfigurationFactory() router = RouterFactory(id=1, name=model_configuration.name, type=RouterType.TEXT_GENERATION) router_repository.get_all_routers.return_value = [] router_repository.create_router.return_value = router - provider_capabilities_repository.get_provider_capabilities.return_value = ProviderNotReachableError( + provider_capabilities_probe.get_capabilities.return_value = ProviderNotReachableError( model_name="my-model", status_code=500, detail="error_detail" ) @@ -268,15 +268,13 @@ async def test_returns_provider_not_reachable_error_and_rolls_back( router_repository.delete_all_routers.assert_awaited_once() @pytest.mark.asyncio - async def test_returns_model_not_found_error_and_rolls_back( - self, use_case, router_repository, provider_repository, provider_capabilities_repository - ): + async def test_returns_model_not_found_error_and_rolls_back(self, use_case, router_repository, provider_repository, provider_capabilities_probe): # Arrange model_configuration = ModelConfigurationFactory() router = RouterFactory(id=1, name=model_configuration.name, type=RouterType.TEXT_GENERATION) router_repository.get_all_routers.return_value = [] router_repository.create_router.return_value = router - provider_capabilities_repository.get_provider_capabilities.return_value = ModelNotFoundError(name="my-model") + provider_capabilities_probe.get_capabilities.return_value = ModelNotFoundError(name="my-model") # Act result = await use_case.execute(routers_to_create=[model_configuration], bootstrap_admin_user_id=BOOTSTRAP_ADMIN_USER_ID) @@ -292,7 +290,7 @@ async def test_returns_inconsistent_max_context_length_error_and_rolls_back( use_case, router_repository, provider_repository, - provider_capabilities_repository, + provider_capabilities_probe, ): # Arrange model_configuration = ModelConfigurationFactory( @@ -304,7 +302,7 @@ async def test_returns_inconsistent_max_context_length_error_and_rolls_back( router = RouterFactory(id=1, name=model_configuration.name, type=RouterType.TEXT_GENERATION, max_context_length=4096, vector_size=None) router_repository.get_all_routers.return_value = [] router_repository.create_router.return_value = router - provider_capabilities_repository.get_provider_capabilities.side_effect = [ + provider_capabilities_probe.get_capabilities.side_effect = [ ProviderCapabilities(max_context_length=4096, vector_size=None), ProviderCapabilities(max_context_length=2048, vector_size=None), ] @@ -322,7 +320,7 @@ async def test_returns_inconsistent_max_context_length_error_and_rolls_back( @pytest.mark.asyncio async def test_returns_inconsistent_vector_size_error_and_rolls_back( - self, use_case, router_repository, provider_repository, provider_capabilities_repository + self, use_case, router_repository, provider_repository, provider_capabilities_probe ): # Arrange model_configuration = ModelConfigurationFactory( @@ -337,7 +335,7 @@ async def test_returns_inconsistent_vector_size_error_and_rolls_back( ) router_repository.get_all_routers.return_value = [] router_repository.create_router.return_value = router - provider_capabilities_repository.get_provider_capabilities.side_effect = [ + provider_capabilities_probe.get_capabilities.side_effect = [ ProviderCapabilities(max_context_length=512, vector_size=768), ProviderCapabilities(max_context_length=512, vector_size=384), ] @@ -359,9 +357,7 @@ async def test_returns_inconsistent_vector_size_error_and_rolls_back( router_repository.delete_all_routers.assert_awaited_once() @pytest.mark.asyncio - async def test_returns_success_with_no_routers_to_create( - self, use_case, router_repository, provider_repository, provider_capabilities_repository - ): + async def test_returns_success_with_no_routers_to_create(self, use_case, router_repository, provider_repository, provider_capabilities_probe): # Arrange router_repository.get_all_routers.return_value = [] @@ -371,7 +367,7 @@ async def test_returns_success_with_no_routers_to_create( # Assert assert result == BootstrapModelsUseCaseSuccess(number_of_routers=0) router_repository.create_router.assert_not_awaited() - provider_capabilities_repository.get_provider_capabilities.assert_not_awaited() + provider_capabilities_probe.get_capabilities.assert_not_awaited() provider_repository.create_provider.assert_not_awaited() diff --git a/api/tests/unit/domain/provider/__init__.py b/api/tests/unit/use_case/services/__init__.py similarity index 100% rename from api/tests/unit/domain/provider/__init__.py rename to api/tests/unit/use_case/services/__init__.py diff --git a/api/tests/unit/domain/provider/test_providercapabilitiesrepository.py b/api/tests/unit/use_case/services/test_providercapabilitiesprobe.py similarity index 74% rename from api/tests/unit/domain/provider/test_providercapabilitiesrepository.py rename to api/tests/unit/use_case/services/test_providercapabilitiesprobe.py index 8f9998c56..83a931297 100644 --- a/api/tests/unit/domain/provider/test_providercapabilitiesrepository.py +++ b/api/tests/unit/use_case/services/test_providercapabilitiesprobe.py @@ -6,7 +6,6 @@ from api.domain.model.entities import Model, Models from api.domain.model.entities import ModelType as RouterType from api.domain.model.errors import ModelNotFoundError, StatusCodeModelError -from api.domain.provider import ProviderCapabilitiesRepository from api.domain.provider.entities import ProviderCapabilities, ProviderFormattedResponse, ProviderOriginalResponse, ProviderType from api.domain.provider.errors import ProviderNotReachableError from api.infrastructure.http import HttpProviderAdapterBuilder @@ -17,6 +16,7 @@ from api.tests.integration.factories.albert import AlbertModelResponseFactory, AlbertModelsResponseFactory from api.tests.integration.factories.tei import TeiEmbeddingsResponseFactory from api.tests.unit.use_case.factories import ProviderFactory +from api.use_cases.services import ProviderCapabilitiesProbe DEFAULT_PROVIDER_URL = "https://test.com" DEFAULT_MODEL_ID = "test-model" @@ -35,8 +35,8 @@ def provider_adapter_builder() -> HttpProviderAdapterBuilder: @pytest.fixture -def repository(provider_client: Mock, provider_adapter_builder: HttpProviderAdapterBuilder) -> ProviderCapabilitiesRepository: - return ProviderCapabilitiesRepository(provider_client=provider_client, provider_adapter_builder=provider_adapter_builder) +def probe(provider_client: Mock, provider_adapter_builder: HttpProviderAdapterBuilder) -> ProviderCapabilitiesProbe: + return ProviderCapabilitiesProbe(provider_client=provider_client, provider_adapter_builder=provider_adapter_builder) def provider_factory(provider_type: ProviderType = ProviderType.ALBERT, model_name: str = DEFAULT_MODEL_ID): @@ -51,13 +51,13 @@ def embeddings_adapter() -> TeiEmbeddingsAdapter: return TeiEmbeddingsAdapter(provider=provider_factory(provider_type=ProviderType.TEI)) -class TestProviderCapabilitiesRepository: +class TestProviderCapabilitiesProbe: @pytest.mark.asyncio - async def test_should_get_capabilities_for_generation_router(self, repository: ProviderCapabilitiesRepository, mocker): - mocked_get_max_context_length = mocker.patch.object(ProviderCapabilitiesRepository, "_get_max_context_length", AsyncMock(return_value=4096)) - mocked_get_vector_size = mocker.patch.object(ProviderCapabilitiesRepository, "_get_vector_size", AsyncMock()) + async def test_should_get_capabilities_for_generation_router(self, probe: ProviderCapabilitiesProbe, mocker): + mocked_get_max_context_length = mocker.patch.object(ProviderCapabilitiesProbe, "_get_max_context_length", AsyncMock(return_value=4096)) + mocked_get_vector_size = mocker.patch.object(ProviderCapabilitiesProbe, "_get_vector_size", AsyncMock()) - result = await repository.get_provider_capabilities( + result = await probe.get_capabilities( router_type=RouterType.TEXT_GENERATION, provider_type=ProviderType.ALBERT, url=DEFAULT_PROVIDER_URL, @@ -75,11 +75,11 @@ async def test_should_get_capabilities_for_generation_router(self, repository: P mocked_get_vector_size.assert_not_called() @pytest.mark.asyncio - async def test_should_get_capabilities_for_embedding_router(self, repository: ProviderCapabilitiesRepository, mocker): - mocked_get_max_context_length = mocker.patch.object(ProviderCapabilitiesRepository, "_get_max_context_length", AsyncMock(return_value=2048)) - mocked_get_vector_size = mocker.patch.object(ProviderCapabilitiesRepository, "_get_vector_size", AsyncMock(return_value=3)) + async def test_should_get_capabilities_for_embedding_router(self, probe: ProviderCapabilitiesProbe, mocker): + mocked_get_max_context_length = mocker.patch.object(ProviderCapabilitiesProbe, "_get_max_context_length", AsyncMock(return_value=2048)) + mocked_get_vector_size = mocker.patch.object(ProviderCapabilitiesProbe, "_get_vector_size", AsyncMock(return_value=3)) - result = await repository.get_provider_capabilities( + result = await probe.get_capabilities( router_type=RouterType.TEXT_EMBEDDINGS_INFERENCE, provider_type=ProviderType.TEI, url=DEFAULT_PROVIDER_URL, @@ -100,10 +100,10 @@ async def test_should_get_capabilities_for_embedding_router(self, repository: Pr "error", [ProviderNotReachableError(model_name=DEFAULT_MODEL_ID, status_code=500, detail="error_detail"), ModelNotFoundError(name=DEFAULT_MODEL_ID)], ) - async def test_should_return_max_context_error(self, repository: ProviderCapabilitiesRepository, error, mocker): - mocker.patch.object(ProviderCapabilitiesRepository, "_get_max_context_length", AsyncMock(return_value=error)) + async def test_should_return_max_context_error(self, probe: ProviderCapabilitiesProbe, error, mocker): + mocker.patch.object(ProviderCapabilitiesProbe, "_get_max_context_length", AsyncMock(return_value=error)) - result = await repository.get_provider_capabilities( + result = await probe.get_capabilities( router_type=RouterType.TEXT_GENERATION, provider_type=ProviderType.ALBERT, url=DEFAULT_PROVIDER_URL, @@ -115,12 +115,12 @@ async def test_should_return_max_context_error(self, repository: ProviderCapabil assert result == error @pytest.mark.asyncio - async def test_should_return_vector_size_error(self, repository: ProviderCapabilitiesRepository, mocker): + async def test_should_return_vector_size_error(self, probe: ProviderCapabilitiesProbe, mocker): error = ProviderNotReachableError(model_name=DEFAULT_MODEL_ID, status_code=500, detail="error_detail") - mocker.patch.object(ProviderCapabilitiesRepository, "_get_max_context_length", AsyncMock(return_value=4096)) - mocker.patch.object(ProviderCapabilitiesRepository, "_get_vector_size", AsyncMock(return_value=error)) + mocker.patch.object(ProviderCapabilitiesProbe, "_get_max_context_length", AsyncMock(return_value=4096)) + mocker.patch.object(ProviderCapabilitiesProbe, "_get_vector_size", AsyncMock(return_value=error)) - result = await repository.get_provider_capabilities( + result = await probe.get_capabilities( router_type=RouterType.TEXT_EMBEDDINGS_INFERENCE, provider_type=ProviderType.TEI, url=DEFAULT_PROVIDER_URL, @@ -132,14 +132,14 @@ async def test_should_return_vector_size_error(self, repository: ProviderCapabil assert result == error @pytest.mark.asyncio - async def test_should_get_max_context_length_when_model_id_is_found(self, repository: ProviderCapabilitiesRepository, provider_client: Mock): + async def test_should_get_max_context_length_when_model_id_is_found(self, probe: ProviderCapabilitiesProbe, provider_client: Mock): body = AlbertModelsResponseFactory( count=2, data=[AlbertModelResponseFactory(model=DEFAULT_MODEL_ID, aliases=["test-model-alias"], max_context_length=10)], ) provider_client.forward_request.return_value = ProviderOriginalResponse(data=body) - result = await repository._get_max_context_length(adapter=models_adapter()) + result = await probe._get_max_context_length(adapter=models_adapter()) assert result == 10 provider_client.forward_request.assert_awaited_once() @@ -148,7 +148,7 @@ async def test_should_get_max_context_length_when_model_id_is_found(self, reposi assert formatted_request.url == f"{DEFAULT_PROVIDER_URL}/v1/models" @pytest.mark.asyncio - async def test_should_get_max_context_length_when_model_alias_is_found(self, repository: ProviderCapabilitiesRepository, provider_client: Mock): + async def test_should_get_max_context_length_when_model_alias_is_found(self, probe: ProviderCapabilitiesProbe, provider_client: Mock): adapter = Mock() adapter.provider = provider_factory(model_name="model-alias") adapter.format_request.return_value = Mock() @@ -176,13 +176,13 @@ async def test_should_get_max_context_length_when_model_alias_is_found(self, rep ) provider_client.forward_request.return_value = ProviderOriginalResponse(data={}) - result = await repository._get_max_context_length(adapter=adapter) + result = await probe._get_max_context_length(adapter=adapter) assert result == 10 @pytest.mark.asyncio async def test_should_return_the_first_model_max_context_length_when_several_models_with_the_same_name_are_found( - self, repository: ProviderCapabilitiesRepository, provider_client: Mock + self, probe: ProviderCapabilitiesProbe, provider_client: Mock ): body = AlbertModelsResponseFactory( data=[ @@ -192,47 +192,43 @@ async def test_should_return_the_first_model_max_context_length_when_several_mod ) provider_client.forward_request.return_value = ProviderOriginalResponse(data=body) - result = await repository._get_max_context_length(adapter=models_adapter()) + result = await probe._get_max_context_length(adapter=models_adapter()) assert result == 10 @pytest.mark.asyncio - async def test_should_return_model_not_found_when_models_response_is_empty( - self, repository: ProviderCapabilitiesRepository, provider_client: Mock - ): + async def test_should_return_model_not_found_when_models_response_is_empty(self, probe: ProviderCapabilitiesProbe, provider_client: Mock): provider_client.forward_request.return_value = ProviderOriginalResponse(data=AlbertModelsResponseFactory(data=[])) - result = await repository._get_max_context_length(adapter=models_adapter()) + result = await probe._get_max_context_length(adapter=models_adapter()) assert result == ModelNotFoundError(name=DEFAULT_MODEL_ID) @pytest.mark.asyncio async def test_should_return_model_not_found_when_model_is_missing_in_models_response( - self, repository: ProviderCapabilitiesRepository, provider_client: Mock + self, probe: ProviderCapabilitiesProbe, provider_client: Mock ): provider_client.forward_request.return_value = ProviderOriginalResponse(data=AlbertModelsResponseFactory(data=[AlbertModelResponseFactory()])) - result = await repository._get_max_context_length(adapter=models_adapter()) + result = await probe._get_max_context_length(adapter=models_adapter()) assert result == ModelNotFoundError(name=DEFAULT_MODEL_ID) @pytest.mark.asyncio - async def test_should_return_provider_not_reachable_when_getting_max_context_fails( - self, repository: ProviderCapabilitiesRepository, provider_client: Mock - ): + async def test_should_return_provider_not_reachable_when_getting_max_context_fails(self, probe: ProviderCapabilitiesProbe, provider_client: Mock): provider_client.forward_request.return_value = StatusCodeModelError(status_code=500, detail="boom") - result = await repository._get_max_context_length(adapter=models_adapter()) + result = await probe._get_max_context_length(adapter=models_adapter()) assert result == ProviderNotReachableError(model_name=DEFAULT_MODEL_ID, status_code=500, detail="boom") @pytest.mark.asyncio - async def test_should_get_vector_size(self, repository: ProviderCapabilitiesRepository, provider_client: Mock): + async def test_should_get_vector_size(self, probe: ProviderCapabilitiesProbe, provider_client: Mock): provider_client.forward_request.return_value = ProviderOriginalResponse( data=TeiEmbeddingsResponseFactory(dimensions=3, model_id=DEFAULT_MODEL_ID) ) - result = await repository._get_vector_size(adapter=embeddings_adapter()) + result = await probe._get_vector_size(adapter=embeddings_adapter()) assert result == 3 provider_client.forward_request.assert_awaited_once() @@ -242,11 +238,9 @@ async def test_should_get_vector_size(self, repository: ProviderCapabilitiesRepo assert formatted_request.body["model"] == DEFAULT_MODEL_ID @pytest.mark.asyncio - async def test_should_return_provider_not_reachable_when_getting_vector_size_fails( - self, repository: ProviderCapabilitiesRepository, provider_client: Mock - ): + async def test_should_return_provider_not_reachable_when_getting_vector_size_fails(self, probe: ProviderCapabilitiesProbe, provider_client: Mock): provider_client.forward_request.return_value = StatusCodeModelError(status_code=500, detail="boom") - result = await repository._get_vector_size(adapter=embeddings_adapter()) + result = await probe._get_vector_size(adapter=embeddings_adapter()) assert result == ProviderNotReachableError(model_name=DEFAULT_MODEL_ID, status_code=500, detail="boom") diff --git a/api/use_cases/admin/providers/_createproviderusecase.py b/api/use_cases/admin/providers/_createproviderusecase.py index 79d1f4ea6..c5319ad1c 100644 --- a/api/use_cases/admin/providers/_createproviderusecase.py +++ b/api/use_cases/admin/providers/_createproviderusecase.py @@ -1,11 +1,12 @@ from dataclasses import dataclass from api.domain.model.errors import InconsistentModelMaxContextLengthError, InconsistentModelVectorSizeError, ModelNotFoundError -from api.domain.provider import ProviderCapabilitiesRepository, ProviderRepository +from api.domain.provider import ProviderRepository from api.domain.provider.entities import BasicAuth, HostingZone, Metric, Provider, ProviderType from api.domain.provider.errors import InvalidProviderTypeError, ProviderAlreadyExistsError, ProviderNotReachableError from api.domain.router import RouterRepository from api.domain.router.errors import RouterNotFoundError +from api.use_cases.services import ProviderCapabilitiesProbe @dataclass @@ -47,11 +48,11 @@ def __init__( self, router_repository: RouterRepository, provider_repository: ProviderRepository, - provider_capabilities_repository: ProviderCapabilitiesRepository, + provider_capabilities_probe: ProviderCapabilitiesProbe, ): self.router_repository = router_repository self.provider_repository = provider_repository - self.provider_capabilities_repository = provider_capabilities_repository + self.provider_capabilities_probe = provider_capabilities_probe async def execute(self, command: CreateProviderCommand) -> CreateProviderUseCaseResult: router = await self.router_repository.get_router_by_id(router_id=command.router_id) @@ -61,7 +62,7 @@ async def execute(self, command: CreateProviderCommand) -> CreateProviderUseCase if not command.provider_type.is_compatible_with(router_type=router.type): return InvalidProviderTypeError(provider_type=command.provider_type.value, router_type=router.type.value) - result = await self.provider_capabilities_repository.get_provider_capabilities( + result = await self.provider_capabilities_probe.get_capabilities( router_type=router.type, provider_type=command.provider_type, url=command.url, diff --git a/api/use_cases/models/_bootstrapmodelsusecase.py b/api/use_cases/models/_bootstrapmodelsusecase.py index 38d37d07c..5ca5acc6a 100644 --- a/api/use_cases/models/_bootstrapmodelsusecase.py +++ b/api/use_cases/models/_bootstrapmodelsusecase.py @@ -3,11 +3,12 @@ import logging from api.domain.model.errors import InconsistentModelMaxContextLengthError, InconsistentModelVectorSizeError, ModelNotFoundError -from api.domain.provider import ProviderCapabilitiesRepository, ProviderRepository +from api.domain.provider import ProviderRepository from api.domain.provider.errors import ProviderAlreadyExistsError, ProviderNotReachableError from api.domain.router import RouterRepository from api.domain.router.errors import RouterNameAlreadyExistsError from api.schemas.core.configuration import Model as ModelConfiguration +from api.use_cases.services import ProviderCapabilitiesProbe logger = logging.getLogger(__name__) @@ -39,11 +40,11 @@ def __init__( self, router_repository: RouterRepository, provider_repository: ProviderRepository, - provider_capabilities_repository: ProviderCapabilitiesRepository, + provider_capabilities_probe: ProviderCapabilitiesProbe, ): self.router_repository = router_repository self.provider_repository = provider_repository - self.provider_capabilities_repository = provider_capabilities_repository + self.provider_capabilities_probe = provider_capabilities_probe async def execute( self, @@ -85,7 +86,7 @@ async def execute( ) for i, provider_to_create in enumerate(router_to_create.providers): - result = await self.provider_capabilities_repository.get_provider_capabilities( + result = await self.provider_capabilities_probe.get_capabilities( router_type=router.type, provider_type=provider_to_create.type, url=provider_to_create.url, diff --git a/api/use_cases/services/__init__.py b/api/use_cases/services/__init__.py new file mode 100644 index 000000000..d0a0e0dc8 --- /dev/null +++ b/api/use_cases/services/__init__.py @@ -0,0 +1,3 @@ +from ._providercapabilitiesprobe import ProviderCapabilitiesProbe + +__all__ = ["ProviderCapabilitiesProbe"] diff --git a/api/domain/provider/_providercapabilitiesrepository.py b/api/use_cases/services/_providercapabilitiesprobe.py similarity index 93% rename from api/domain/provider/_providercapabilitiesrepository.py rename to api/use_cases/services/_providercapabilitiesprobe.py index 69e8f67a5..3e19db568 100644 --- a/api/domain/provider/_providercapabilitiesrepository.py +++ b/api/use_cases/services/_providercapabilitiesprobe.py @@ -1,20 +1,18 @@ from api.domain.embeddings.entities import CreateEmbeddingsBody from api.domain.model.entities import ModelType as RouterType from api.domain.model.errors import ModelNotFoundError -from api.domain.provider._provideradapter import ProviderAdapter -from api.domain.provider._provideradapterbuilder import ProviderAdapterBuilder -from api.domain.provider._providerclient import ProviderClient +from api.domain.provider import ProviderAdapter, ProviderAdapterBuilder, ProviderClient from api.domain.provider.entities import Provider, ProviderCapabilities, ProviderOriginalRequest, ProviderOriginalResponse, ProviderType from api.domain.provider.errors import ProviderNotReachableError from api.utils.variables import EndpointRoute -class ProviderCapabilitiesRepository: +class ProviderCapabilitiesProbe: def __init__(self, provider_client: ProviderClient, provider_adapter_builder: ProviderAdapterBuilder): self.provider_client = provider_client self.provider_adapter_builder = provider_adapter_builder - async def get_provider_capabilities( + async def get_capabilities( self, router_type: RouterType, provider_type: ProviderType, diff --git a/api/utils/lifespan.py b/api/utils/lifespan.py index bb20f0f72..efcd007c3 100755 --- a/api/utils/lifespan.py +++ b/api/utils/lifespan.py @@ -9,7 +9,6 @@ from api.dependencies import get_postgres_session from api.domain.model.errors import InconsistentModelMaxContextLengthError, InconsistentModelVectorSizeError, ModelNotFoundError -from api.domain.provider import ProviderCapabilitiesRepository from api.domain.provider.errors import ProviderAlreadyExistsError, ProviderNotReachableError from api.domain.router.errors import RouterNameAlreadyExistsError from api.helpers._identityaccessmanager import IdentityAccessManager @@ -36,6 +35,7 @@ BootstrapAdminUseCaseSuccess, ) from api.use_cases.models import BootstrapModelsUseCase, BootstrapModelsUseCaseSkipped, BootstrapModelsUseCaseSuccess +from api.use_cases.services import ProviderCapabilitiesProbe from api.utils.configuration import get_configuration from api.utils.context import global_context from api.utils.logging import init_logger @@ -122,7 +122,7 @@ async def bootstrap_admin_role_and_user(configuration: Configuration, postgres_s async def bootstrap_models(configuration: Configuration, postgres_session: AsyncSession, bootstrap_admin_user_id: int) -> int: router_repository = PostgresRouterRepository(postgres_session=postgres_session, app_title=configuration.settings.app_title) provider_repository = PostgresProviderRepository(postgres_session=postgres_session) - provider_capabilities_repository = ProviderCapabilitiesRepository( + provider_capabilities_probe = ProviderCapabilitiesProbe( provider_client=HttpProviderClient(), provider_adapter_builder=HttpProviderAdapterBuilder(), ) @@ -130,7 +130,7 @@ async def bootstrap_models(configuration: Configuration, postgres_session: Async result = await BootstrapModelsUseCase( router_repository=router_repository, provider_repository=provider_repository, - provider_capabilities_repository=provider_capabilities_repository, + provider_capabilities_probe=provider_capabilities_probe, ).execute(routers_to_create=configuration.models, bootstrap_admin_user_id=bootstrap_admin_user_id) match result: