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 diff --git a/adr/2026-05-28-refactoring-model-forwarding.md b/adr/2026-05-28-refactoring-model-forwarding.md index 16274e3da..4b581f0d6 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 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. @@ -134,7 +135,6 @@ subgraph DL[**Domain layer**] provider_adapter_builder[ProviderAdapterBuilder] provider_adapter[ProviderAdapter] provider_repository[ProviderRepository] - provider_gateway[ProviderGateway] provider_load_balancer[ProviderLoadBalancer] provider_client[ProviderClient] provider_metrics_logger[ProviderMetricsLogger] @@ -187,7 +187,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 @@ -198,8 +197,6 @@ use_case --> usage_computer usage_computer --> model_environmental_impacts_computer usage_computer --> model_tokenizer - - provider_adapter_builder --> http_provider_adapter_builder provider_client --> http_provider_client http_provider_adapter_builder --VLLM, Mistral, TEI...
Chat completions, OCR, Rerank...--> http_provider_adapter diff --git a/api/dependencies.py b/api/dependencies.py index 27a9bb3e8..7d7e57321 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, @@ -56,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 @@ -122,13 +121,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) @@ -170,6 +162,13 @@ def _provider_repository(session: AsyncSession) -> ProviderRepository: return PostgresProviderRepository(postgres_session=session) +def _provider_capabilities_probe( + provider_client: ProviderClient = Depends(_provider_client), + provider_adapter_builder: ProviderAdapterBuilder = Depends(_provider_adapter_builder), +) -> ProviderCapabilitiesProbe: + return ProviderCapabilitiesProbe(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), @@ -362,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_client: ProviderClient = Depends(_provider_client), + provider_capabilities_probe: ProviderCapabilitiesProbe = Depends(_provider_capabilities_probe), ) -> 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_capabilities_probe=provider_capabilities_probe, ) diff --git a/api/domain/provider/__init__.py b/api/domain/provider/__init__.py index a976b5422..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, ProviderGateway from api.domain.provider._providerloadbalancer import ProviderLoadBalancer from api.domain.provider._providermetricslogger import ProviderMetricsLogger from api.domain.provider._providerrepository import ProviderRepository @@ -11,8 +10,6 @@ "ProviderAdapterBuilder", "ProviderClient", "ProviderClientResponse", - "ProviderCapabilities", - "ProviderGateway", "ProviderLoadBalancer", "ProviderMetricsLogger", "ProviderRepository", diff --git a/api/domain/provider/_providergateway.py b/api/domain/provider/_providergateway.py deleted file mode 100644 index f780e92c7..000000000 --- a/api/domain/provider/_providergateway.py +++ /dev/null @@ -1,27 +0,0 @@ -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 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/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/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/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/models/test_modelprovidergateway.py b/api/tests/integration/use_case/test_providercapabilitiesprobe.py similarity index 84% rename from api/tests/integration/models/test_modelprovidergateway.py rename to api/tests/integration/use_case/test_providercapabilitiesprobe.py index 12bad0d6e..f7be808ac 100644 --- a/api/tests/integration/models/test_modelprovidergateway.py +++ b/api/tests/integration/use_case/test_providercapabilitiesprobe.py @@ -5,12 +5,11 @@ 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.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 +from api.use_cases.services import ProviderCapabilitiesProbe DEFAULT_PROVIDER_URL = "http://my-test-provider/" DEFAULT_MODEL_ID = "test/my-model" @@ -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 probe() -> ProviderCapabilitiesProbe: + return ProviderCapabilitiesProbe(provider_client=HttpProviderClient(), provider_adapter_builder=HttpProviderAdapterBuilder()) @pytest.mark.asyncio(loop_scope="session") -class TestModelProviderGateway: +class TestProviderCapabilitiesProbe: @respx.mock - async def test_get_capabilities_of_non_embeddings_providers(self, gateway: ModelProviderGateway): + async def test_get_capabilities_of_non_embeddings_providers(self, probe: ProviderCapabilitiesProbe): _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 probe.get_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, + probe: ProviderCapabilitiesProbe, ): _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 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 e9c148e20..d8e2e6388 100644 --- a/api/tests/unit/use_case/admin/providers/test_createproviderusecase.py +++ b/api/tests/unit/use_case/admin/providers/test_createproviderusecase.py @@ -3,9 +3,8 @@ import pytest 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.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 from api.tests.unit.use_case.factories import ProviderFactory, RouterFactory @@ -23,16 +22,16 @@ def provider_repository(): @pytest.fixture -def provider_gateway(): +def provider_capabilities_probe(): return AsyncMock() @pytest.fixture -def use_case(router_repository, provider_repository, provider_gateway): +def use_case(router_repository, provider_repository, provider_capabilities_probe): return CreateProviderUseCase( router_repository=router_repository, provider_repository=provider_repository, - provider_gateway=provider_gateway, + provider_capabilities_probe=provider_capabilities_probe, ) @@ -169,7 +168,7 @@ async def test_should_create_provider_when_router_has_a_different_provider( use_case, router_repository, provider_repository, - provider_gateway, + provider_capabilities_probe, sample_router_with_providers, sample_provider, default_command, @@ -177,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_gateway.get_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 @@ -210,7 +209,7 @@ async def test_should_create_embedding_provider_when_vector_size_matches( use_case, router_repository, provider_repository, - provider_gateway, + provider_capabilities_probe, sample_embedding_router_with_providers, sample_provider, default_command, @@ -218,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_gateway.get_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 @@ -247,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_gateway, default_command + self, use_case, router_repository, provider_repository, provider_capabilities_probe, default_command ): # Arrange @@ -259,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_gateway.get_capabilities.assert_not_called() + provider_capabilities_probe.get_capabilities.assert_not_called() provider_repository.create_provider.assert_not_called() @pytest.mark.asyncio @@ -273,7 +272,7 @@ async def test_should_create_provider_when_provider_type_is_compatible( use_case, router_repository, provider_repository, - provider_gateway, + provider_capabilities_probe, default_command, router_type, provider_type, @@ -282,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_gateway.get_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) @@ -292,7 +291,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_probe.get_capabilities.assert_called_once_with( router_type=router_type, provider_type=provider_type, url="https://example.com/", @@ -329,7 +328,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_probe, default_command, router_type, provider_type, @@ -345,17 +344,19 @@ 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_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_gateway, 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_gateway.get_capabilities.return_value = ProviderNotReachableError(model_name="my-model", status_code=500, detail="error_detail") + provider_capabilities_probe.get_capabilities.return_value = ProviderNotReachableError( + model_name="my-model", status_code=500, detail="error_detail" + ) # Act result = await use_case.execute(default_command) @@ -367,14 +368,31 @@ 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_probe, sample_router, default_command + ): + # Arrange + + router_repository.get_router_by_id.return_value = sample_router + provider_capabilities_probe.get_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_gateway, 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_gateway.get_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) @@ -391,14 +409,14 @@ async def test_should_return_inconsistent_vector_size_error_when_mismatch( use_case, router_repository, provider_repository, - provider_gateway, + 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_gateway.get_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)) @@ -411,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_gateway, 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_gateway.get_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 cc2250aa4..37fb9abd8 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 ( @@ -29,18 +29,22 @@ def provider_repository(): @pytest.fixture -def provider_gateway(): +def provider_capabilities_probe(): 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_probe): + return BootstrapModelsUseCase( + router_repository=router_repository, + provider_repository=provider_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_gateway): + 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 @@ -51,11 +55,13 @@ 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_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_gateway): + async def test_successfully_creates_router_with_single_provider( + self, use_case, router_repository, provider_repository, provider_capabilities_probe + ): # Arrange model_provider = ModelProviderConfigurationFactory() model_configuration = ModelConfigurationFactory(providers=[model_provider]) @@ -64,7 +70,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_probe.get_capabilities.return_value = ProviderCapabilities(max_context_length=4096, vector_size=None) provider_repository.create_provider.return_value = provider # Act @@ -82,7 +88,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_probe.get_capabilities.assert_awaited_once_with( router_type=router.type, provider_type=model_provider.type, url=model_provider.url, @@ -111,7 +117,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_probe ): # Arrange first_model = ModelConfigurationFactory( @@ -133,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_gateway.get_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), @@ -153,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_gateway.get_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_gateway + self, use_case, router_repository, provider_repository, provider_capabilities_probe ): # Arrange routers_to_create = [ @@ -177,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_gateway.get_capabilities.assert_not_awaited() + provider_capabilities_probe.get_capabilities.assert_not_awaited() provider_repository.create_provider.assert_not_awaited() @pytest.mark.asyncio @@ -218,7 +224,7 @@ async def test_returns_provider_already_exists_error_when_duplicate_within_route use_case, router_repository, provider_repository, - provider_gateway, + provider_capabilities_probe, ): # Arrange model_configuration = ModelConfigurationFactory( @@ -237,17 +243,21 @@ 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_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_gateway): + async def test_returns_provider_not_reachable_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_gateway.get_capabilities.return_value = ProviderNotReachableError(model_name="my-model", status_code=500, detail="error_detail") + provider_capabilities_probe.get_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 +268,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_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_gateway.get_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) @@ -280,7 +290,7 @@ async def test_returns_inconsistent_max_context_length_error_and_rolls_back( use_case, router_repository, provider_repository, - provider_gateway, + provider_capabilities_probe, ): # Arrange model_configuration = ModelConfigurationFactory( @@ -292,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_gateway.get_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), ] @@ -309,7 +319,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_gateway): + async def test_returns_inconsistent_vector_size_error_and_rolls_back( + self, use_case, router_repository, provider_repository, provider_capabilities_probe + ): # Arrange model_configuration = ModelConfigurationFactory( type=RouterType.TEXT_EMBEDDINGS_INFERENCE, @@ -323,7 +335,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_probe.get_capabilities.side_effect = [ ProviderCapabilities(max_context_length=512, vector_size=768), ProviderCapabilities(max_context_length=512, vector_size=384), ] @@ -345,7 +357,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_probe): # Arrange router_repository.get_all_routers.return_value = [] @@ -355,7 +367,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_probe.get_capabilities.assert_not_awaited() provider_repository.create_provider.assert_not_awaited() diff --git a/api/tests/unit/use_case/services/__init__.py b/api/tests/unit/use_case/services/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/api/tests/unit/infrastructure/model/test_modelprovidergateway.py b/api/tests/unit/use_case/services/test_providercapabilitiesprobe.py similarity index 73% rename from api/tests/unit/infrastructure/model/test_modelprovidergateway.py rename to api/tests/unit/use_case/services/test_providercapabilitiesprobe.py index d3f8a4874..83a931297 100644 --- a/api/tests/unit/infrastructure/model/test_modelprovidergateway.py +++ b/api/tests/unit/use_case/services/test_providercapabilitiesprobe.py @@ -6,18 +6,17 @@ 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.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 +from api.use_cases.services import ProviderCapabilitiesProbe DEFAULT_PROVIDER_URL = "https://test.com" DEFAULT_MODEL_ID = "test-model" @@ -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 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): @@ -52,13 +51,13 @@ def embeddings_adapter() -> TeiEmbeddingsAdapter: return TeiEmbeddingsAdapter(provider=provider_factory(provider_type=ProviderType.TEI)) -class TestModelProviderGateway: +class TestProviderCapabilitiesProbe: @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, 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 gateway.get_capabilities( + result = await probe.get_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, 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 gateway.get_capabilities( + result = await probe.get_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, probe: ProviderCapabilitiesProbe, error, mocker): + mocker.patch.object(ProviderCapabilitiesProbe, "_get_max_context_length", AsyncMock(return_value=error)) - result = await gateway.get_capabilities( + result = await probe.get_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, probe: ProviderCapabilitiesProbe, 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(ProviderCapabilitiesProbe, "_get_max_context_length", AsyncMock(return_value=4096)) + mocker.patch.object(ProviderCapabilitiesProbe, "_get_vector_size", AsyncMock(return_value=error)) - result = await gateway.get_capabilities( + result = await probe.get_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, 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 gateway._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() @@ -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, probe: ProviderCapabilitiesProbe, 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 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, gateway: ModelProviderGateway, provider_client: Mock + self, probe: ProviderCapabilitiesProbe, provider_client: Mock ): body = AlbertModelsResponseFactory( data=[ @@ -193,41 +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 gateway._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, gateway: ModelProviderGateway, 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 gateway._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, gateway: ModelProviderGateway, provider_client: Mock): + async def test_should_return_model_not_found_when_model_is_missing_in_models_response( + self, probe: ProviderCapabilitiesProbe, 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 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, gateway: ModelProviderGateway, 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 gateway._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, gateway: ModelProviderGateway, 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 gateway._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() @@ -237,9 +238,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, probe: ProviderCapabilitiesProbe, 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 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 bc8e2e09c..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 -from api.domain.provider import ProviderGateway, ProviderRepository +from api.domain.model.errors import InconsistentModelMaxContextLengthError, InconsistentModelVectorSizeError, ModelNotFoundError +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 @@ -34,6 +35,7 @@ class CreateProviderUseCaseSuccess: CreateProviderUseCaseSuccess | InvalidProviderTypeError | ProviderNotReachableError + | ModelNotFoundError | InconsistentModelMaxContextLengthError | InconsistentModelVectorSizeError | RouterNotFoundError @@ -42,10 +44,15 @@ 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_capabilities_probe: ProviderCapabilitiesProbe, + ): self.router_repository = router_repository self.provider_repository = provider_repository - self.provider_gateway = provider_gateway + 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) @@ -55,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_gateway.get_capabilities( + result = await self.provider_capabilities_probe.get_capabilities( router_type=router.type, provider_type=command.provider_type, url=command.url, @@ -66,6 +73,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 diff --git a/api/use_cases/models/_bootstrapmodelsusecase.py b/api/use_cases/models/_bootstrapmodelsusecase.py index 7b32b2e1c..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 ProviderGateway, 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__) @@ -35,10 +36,15 @@ 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_capabilities_probe: ProviderCapabilitiesProbe, + ): self.router_repository = router_repository self.provider_repository = provider_repository - self.provider_gateway = provider_gateway + self.provider_capabilities_probe = provider_capabilities_probe async def execute( self, @@ -80,7 +86,7 @@ async def execute( ) for i, provider_to_create in enumerate(router_to_create.providers): - result = await self.provider_gateway.get_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/infrastructure/model/_modelprovidergateway.py b/api/use_cases/services/_providercapabilitiesprobe.py similarity index 79% rename from api/infrastructure/model/_modelprovidergateway.py rename to api/use_cases/services/_providercapabilitiesprobe.py index b70b23c4e..3e19db568 100644 --- a/api/infrastructure/model/_modelprovidergateway.py +++ b/api/use_cases/services/_providercapabilitiesprobe.py @@ -1,22 +1,16 @@ -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 import ProviderAdapter, ProviderAdapterBuilder, ProviderClient +from api.domain.provider.entities import Provider, ProviderCapabilities, 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): +class ProviderCapabilitiesProbe: def __init__(self, provider_client: ProviderClient, provider_adapter_builder: ProviderAdapterBuilder): + self.provider_client = provider_client self.provider_adapter_builder = provider_adapter_builder - self.client = provider_client async def get_capabilities( self, @@ -62,10 +56,10 @@ async def get_capabilities( 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: + 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 self.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 @@ -80,13 +74,13 @@ async def _get_max_context_length(self, adapter: ModelsAdapter) -> int | None | return model.max_context_length - async def _get_vector_size(self, adapter: EmbeddingsAdapter) -> 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 self.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/utils/lifespan.py b/api/utils/lifespan.py index b100e43db..efcd007c3 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, @@ -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,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_gateway = ModelProviderGateway(provider_client=provider_client, provider_adapter_builder=provider_adapter_builder) + provider_capabilities_probe = ProviderCapabilitiesProbe( + provider_client=HttpProviderClient(), + provider_adapter_builder=HttpProviderAdapterBuilder(), + ) result = await BootstrapModelsUseCase( router_repository=router_repository, provider_repository=provider_repository, - provider_gateway=provider_gateway, + provider_capabilities_probe=provider_capabilities_probe, ).execute(routers_to_create=configuration.models, bootstrap_admin_user_id=bootstrap_admin_user_id) match result: