diff --git a/.github/badges/coverage.json b/.github/badges/coverage.json index 68cd04f7e..fee9ce1f1 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.83%", "color": "red"} \ No newline at end of file diff --git a/api/domain/embeddings/entities.py b/api/domain/embeddings/entities.py index b72a85c1b..afb30e5d7 100644 --- a/api/domain/embeddings/entities.py +++ b/api/domain/embeddings/entities.py @@ -4,19 +4,26 @@ from typing import Any, Literal from openai.types import CreateEmbeddingResponse +from openai.types.chat import ChatCompletionContentPartParam from pydantic import Field from api.domain import BaseModel from api.domain.usage.entities import Usage +class EmbeddingMessage(BaseModel): + role: Literal["system", "user", "assistant", "developer", "function", "tool"] + content: str | list[ChatCompletionContentPartParam] + + class EncodingFormat(StrEnum): FLOAT = "float" BASE64 = "base64" class CreateEmbeddingsBody(BaseModel): - input: list[int] | list[list[int]] | str | list[str] + input: list[int] | list[list[int]] | str | list[str] | None + messages: list[EmbeddingMessage] | None = None model: str dimensions: int | None = None encoding_format: EncodingFormat = EncodingFormat.FLOAT @@ -24,11 +31,23 @@ class CreateEmbeddingsBody(BaseModel): def get_prompts(self) -> list[str]: if isinstance(self.input, str): return [self.input] - elif isinstance(self.input, list): + elif isinstance(self.input, list) and len(self.input) > 0: if isinstance(self.input[0], list): return [str(item) for sublist in self.input for item in sublist] else: return [str(item) for item in self.input] + elif isinstance(self.messages, list) and len(self.messages) > 0: + prompts = [] + for message in self.messages: + if isinstance(message.content, list): + for content_part in message.content: + if content_part["type"] == "text": + prompts.append(content_part["text"]) + else: + prompts.append(message.content) + return prompts + else: + return [] class Embeddings(CreateEmbeddingResponse): diff --git a/api/infrastructure/fastapi/endpoints/embeddings.py b/api/infrastructure/fastapi/endpoints/embeddings.py index 16bf2f49c..6fb4a088a 100644 --- a/api/infrastructure/fastapi/endpoints/embeddings.py +++ b/api/infrastructure/fastapi/endpoints/embeddings.py @@ -53,13 +53,7 @@ async def create_embeddings( request_context: ContextVar[RequestContext] = Depends(get_request_context), ) -> JSONResponse: try: - command = CreateEmbeddingsCommand( - input=body.input, - model=body.model, - dimensions=body.dimensions, - encoding_format=body.encoding_format, - request_context=request_context, - ) + command = CreateEmbeddingsCommand(**body.model_dump(), request_context=request_context) result = await create_embeddings_use_case.execute(command) except Exception as e: logger.exception( diff --git a/api/infrastructure/fastapi/endpoints/rerank.py b/api/infrastructure/fastapi/endpoints/rerank.py index deda65b39..34370d2da 100644 --- a/api/infrastructure/fastapi/endpoints/rerank.py +++ b/api/infrastructure/fastapi/endpoints/rerank.py @@ -53,13 +53,7 @@ async def create_rerank( request_context: ContextVar[RequestContext] = Depends(get_request_context), ) -> JSONResponse: try: - command = CreateRerankCommand( - query=body.query, - documents=body.documents, - model=body.model, - top_n=body.top_n, - request_context=request_context, - ) + command = CreateRerankCommand(**body.model_dump(), request_context=request_context) result = await create_rerank_use_case.execute(command) except Exception as e: logger.exception( diff --git a/api/infrastructure/fastapi/schemas/embeddings.py b/api/infrastructure/fastapi/schemas/embeddings.py index 82041751b..a4a3c7c55 100644 --- a/api/infrastructure/fastapi/schemas/embeddings.py +++ b/api/infrastructure/fastapi/schemas/embeddings.py @@ -2,7 +2,9 @@ from typing import Annotated, Literal from openai.types import CreateEmbeddingResponse +from openai.types.chat import ChatCompletionContentPartParam from pydantic import Field, StringConstraints +from pydantic.json_schema import SkipJsonSchema from api.domain import BaseModel from api.domain.usage.entities import Usage @@ -13,9 +15,15 @@ class EncodingFormat(StrEnum): BASE64 = "base64" +class EmbeddingMessage(BaseModel): + role: Annotated[Literal["system", "user", "assistant", "developer", "function", "tool"], Field(description="The role of the message.")] + content: Annotated[str | list[ChatCompletionContentPartParam], Field(description="The content of the message.")] + + class CreateEmbeddingsBody(BaseModel): - input: Annotated[list[int] | list[Annotated[list[int], Field(min_length=1)]] | str | list[str], Field(min_length=1)] = Field(default=..., description="Input text to embed, encoded as a string or array of tokens. To embed multiple inputs in a single request, pass an array of strings or array of token arrays. The input must not exceed the max input tokens for the model (call `/v1/models` endpoint to get the `max_context_length` by model) and cannot be an empty string.") # fmt: off model: Annotated[str, StringConstraints(min_length=1), Field(default=..., description="ID of the model to use. Call `/v1/models` endpoint to get the list of available models, only `text-embeddings-inference` model type is supported.")] # fmt: off + input: Annotated[list[int] | list[Annotated[list[int], Field(min_length=1)]] | str | list[str], Field(min_length=1)] | None = Field(default=None, description="Input text to embed, encoded as a string or array of tokens. To embed multiple inputs in a single request, pass an array of strings or array of token arrays. The input must not exceed the max input tokens for the model (call `/v1/models` endpoint to get the `max_context_length` by model) and cannot be an empty string.") # fmt: off + messages: Annotated[SkipJsonSchema[list[EmbeddingMessage]] | None, Field(default=None)] dimensions: Annotated[int | None, Field(default=None, gt=0, description="The number of dimensions the resulting output embeddings should have.")] # fmt: off encoding_format: Annotated[EncodingFormat, Field(default=EncodingFormat.FLOAT, description="The format of the output embeddings.")] # fmt: off diff --git a/api/infrastructure/http/adapters/_httpprovideradapter.py b/api/infrastructure/http/adapters/_httpprovideradapter.py index 76960fd65..0f81c3a32 100644 --- a/api/infrastructure/http/adapters/_httpprovideradapter.py +++ b/api/infrastructure/http/adapters/_httpprovideradapter.py @@ -31,7 +31,7 @@ def format_request(self, original_request: ProviderOriginalRequest) -> ProviderF formatted_request = ProviderFormattedRequest( method=self.TARGET_ENDPOINT_METHOD, url=target_url, - body=original_request.body.model_dump() if original_request.body else {}, + body=original_request.body.model_dump(exclude_none=True) if original_request.body else {}, form=original_request.form if original_request.form else {}, files=original_request.files if original_request.files else {}, ) diff --git a/api/infrastructure/http/adapters/chat/mistral/_mistralchatcompletionsadapter.py b/api/infrastructure/http/adapters/chat/mistral/_mistralchatcompletionsadapter.py index b2c7a5c91..5f06f52d1 100644 --- a/api/infrastructure/http/adapters/chat/mistral/_mistralchatcompletionsadapter.py +++ b/api/infrastructure/http/adapters/chat/mistral/_mistralchatcompletionsadapter.py @@ -4,28 +4,29 @@ class MistralChatCompletionsAdapter(ChatCompletionsAdapter): def format_request(self, original_request: ProviderOriginalRequest) -> ProviderFormattedRequest: - # @TODO: build body with model_fields_set to exclude unset fields - # see https://docs.mistral.ai/api#operation-chat_completion_v1_chat_completions_post - body = { - "frequency_penalty": original_request.body.frequency_penalty or 0.0, - "max_tokens": original_request.body.max_tokens, - "messages": original_request.body.messages, - "model": self.provider.model_name, - "n": original_request.body.n, - "parallel_tool_calls": original_request.body.parallel_tool_calls or False, - "prediction": original_request.body.prediction or {}, - "presence_penalty": original_request.body.presence_penalty or 0.0, - "prompt_mode": original_request.body.prompt_mode, - "random_seed": original_request.body.random_seed or original_request.body.seed, - "response_format": original_request.body.response_format or {"type": "text"}, - "safe_prompt": original_request.body.safe_prompt or False, - "stop": original_request.body.stop or [], - "stream": original_request.body.stream or False, - "temperature": original_request.body.temperature, - "tool_choice": original_request.body.tool_choice, - "tools": original_request.body.tools, - "top_p": original_request.body.top_p or 1.0, - } + body = original_request.body.model_dump(exclude_none=True) + body["random_seed"] = body["random_seed"] or body["seed"] + supported_fields = [ + "frequency_penalty", + "max_tokens", + "messages", + "model", + "n", + "parallel_tool_calls", + "prediction", + "presence_penalty", + "prompt_mode", + "random_seed", + "response_format", + "safe_prompt", + "stop", + "stream", + "temperature", + "tool_choice", + "tools", + "top_p", + ] + body = {key: value for key, value in body.items() if key in supported_fields} target_url = self._build_target_url(base_url=self.provider.url, target_endpoint_route=self.TARGET_ENDPOINT_ROUTE) return ProviderFormattedRequest(method=self.TARGET_ENDPOINT_METHOD, url=target_url, body=body) diff --git a/api/infrastructure/http/adapters/rerank/tei/_teireranksadapter.py b/api/infrastructure/http/adapters/rerank/tei/_teireranksadapter.py index 76de38819..6f90da082 100644 --- a/api/infrastructure/http/adapters/rerank/tei/_teireranksadapter.py +++ b/api/infrastructure/http/adapters/rerank/tei/_teireranksadapter.py @@ -27,7 +27,7 @@ def format_request(self, original_request: ProviderOriginalRequest) -> ProviderF { "query": original_request.body.query, "texts": original_request.body.documents, - **original_request.body.model_dump(), + **original_request.body.model_dump(exclude_none=True), } ) except ValidationError as e: diff --git a/api/tests/integration/endpoints/test_embeddings.py b/api/tests/integration/endpoints/test_embeddings.py index 5e5a5b4a6..6197a70e5 100644 --- a/api/tests/integration/endpoints/test_embeddings.py +++ b/api/tests/integration/endpoints/test_embeddings.py @@ -1,3 +1,4 @@ +import json from unittest.mock import AsyncMock, MagicMock from httpx import AsyncClient @@ -90,6 +91,37 @@ async def test_happy_path(self, client: AsyncClient, db_session): assert len(data["data"]) >= 1 assert all("embedding" in item and "index" in item for item in data["data"]) + @respx.mock + async def test_omitted_optional_fields_are_excluded_from_provider_body(self, client: AsyncClient, db_session): + admin_key = await create_key(db_session, name="admin_embeddings_exclude_none_key", user=self.router_owner) + RouterSQLFactory( + user=self.router_owner, + name=DEFAULT_MODEL_NAME, + type=ModelType.TEXT_EMBEDDINGS_INFERENCE, + providers=1, + providers__type=ProviderType.TEI, + providers__url=DEFAULT_PROVIDER_URL, + ) + await db_session.flush() + + route = mock_embeddings_responses( + respx_mock=respx, + provider_type=ProviderType.TEI, + body=TeiEmbeddingsResponseFactory(), + status_code=TeiEmbeddingsResponseFactory._status_code, + ) + + response = await client.post( + url=URL, + headers={"Authorization": f"Bearer {admin_key.token}"}, + json=_valid_body(), + ) + + assert response.status_code == 200, response.text + provider_json = json.loads(route.calls[0].request.content) + assert "dimensions" not in provider_json + assert None not in provider_json.values() + @pytest.mark.parametrize( "use_case_result,expected_status,expected_detail", [ diff --git a/api/tests/integration/endpoints/test_rerank.py b/api/tests/integration/endpoints/test_rerank.py index 997933304..0e9e46af5 100644 --- a/api/tests/integration/endpoints/test_rerank.py +++ b/api/tests/integration/endpoints/test_rerank.py @@ -1,3 +1,4 @@ +import json from unittest.mock import AsyncMock, MagicMock from httpx import AsyncClient @@ -96,6 +97,37 @@ async def test_happy_path(self, client: AsyncClient, db_session): assert len(data["results"]) == len(DEFAULT_DOCUMENTS) assert all("relevance_score" in result and "index" in result for result in data["results"]) + @respx.mock + async def test_omitted_optional_fields_are_excluded_from_provider_body(self, client: AsyncClient, db_session): + admin_key = await create_key(db_session, name="admin_rerank_exclude_none_key", user=self.router_owner) + RouterSQLFactory( + user=self.router_owner, + name=DEFAULT_MODEL_NAME, + type=ModelType.TEXT_CLASSIFICATION, + providers=1, + providers__type=ProviderType.TEI, + providers__url=DEFAULT_PROVIDER_URL, + ) + await db_session.flush() + + route = mock_rerank_responses( + respx_mock=respx, + provider_type=ProviderType.TEI, + body=TeiRerankResponseFactory(count=len(DEFAULT_DOCUMENTS)), + status_code=TeiRerankResponseFactory._status_code, + ) + + response = await client.post( + url=URL, + headers={"Authorization": f"Bearer {admin_key.token}"}, + json=_valid_body(), + ) + + assert response.status_code == 200, response.text + provider_json = json.loads(route.calls[0].request.content) + assert "top_n" not in provider_json + assert None not in provider_json.values() + @pytest.mark.parametrize( "use_case_result,expected_status,expected_detail", [ diff --git a/api/tests/integration/endpoints/utils.py b/api/tests/integration/endpoints/utils.py index b35d852fe..2beacd934 100644 --- a/api/tests/integration/endpoints/utils.py +++ b/api/tests/integration/endpoints/utils.py @@ -36,14 +36,14 @@ def mock_models_responses(respx_mock, provider_type: ProviderType, body: factory respx_mock.get(url=url).mock(return_value=httpx.Response(status_code=status_code, json=body)) -def mock_embeddings_responses(respx_mock, provider_type: ProviderType, body: factory.DictFactory, status_code: int) -> None: +def mock_embeddings_responses(respx_mock, provider_type: ProviderType, body: factory.DictFactory, status_code: int): url = urljoin(DEFAULT_PROVIDER_URL, url=EMBEDDINGS_ENDPOINT_BY_PROVIDER[provider_type]) - respx_mock.post(url=url).mock(return_value=httpx.Response(status_code=status_code, json=body)) + return respx_mock.post(url=url).mock(return_value=httpx.Response(status_code=status_code, json=body)) -def mock_rerank_responses(respx_mock, provider_type: ProviderType, body: list | factory.Factory, status_code: int) -> None: +def mock_rerank_responses(respx_mock, provider_type: ProviderType, body: list | factory.Factory, status_code: int): url = urljoin(DEFAULT_PROVIDER_URL, RERANK_ENDPOINT_BY_PROVIDER[provider_type]) - respx_mock.post(url=url).mock(return_value=httpx.Response(status_code=status_code, json=body)) + return respx_mock.post(url=url).mock(return_value=httpx.Response(status_code=status_code, json=body)) def mock_metrics_responses(respx_mock, provider_type: ProviderType, text: str, status_code: int) -> None: diff --git a/api/tests/unit/domain/embeddings/test_embeddingsentities.py b/api/tests/unit/domain/embeddings/test_embeddingsentities.py index f739850fa..60674e5b8 100644 --- a/api/tests/unit/domain/embeddings/test_embeddingsentities.py +++ b/api/tests/unit/domain/embeddings/test_embeddingsentities.py @@ -48,3 +48,58 @@ def test_returns_flattened_string_items_when_input_is_a_list_of_lists_of_integer # Assert assert result == ["1", "2", "3", "4", "5", "6"] + + def test_returns_messages_content_when_input_is_none(self, embeddings_body: CreateEmbeddingsBody): + # Arrange + embeddings_body = CreateEmbeddingsBody( + model=embeddings_body.model, + input=None, + messages=[ + {"role": "system", "content": "system prompt"}, + {"role": "user", "content": "user prompt"}, + ], + ) + + # Act + result = embeddings_body.get_prompts() + + # Assert + assert result == ["system prompt", "user prompt"] + + def test_returns_only_text_parts_from_multimodal_messages_when_input_is_none(self, embeddings_body: CreateEmbeddingsBody): + # Arrange + embeddings_body = CreateEmbeddingsBody( + model=embeddings_body.model, + input=None, + messages=[ + { + "role": "user", + "content": [ + {"type": "image_url", "image_url": {"url": "https://example.com/image.jpg"}}, + {"type": "text", "text": "Describe this image"}, + ], + }, + { + "role": "assistant", + "content": [ + {"type": "text", "text": "Image description"}, + ], + }, + ], + ) + + # Act + result = embeddings_body.get_prompts() + + # Assert + assert result == ["Describe this image", "Image description"] + + def test_returns_empty_list_when_input_and_messages_are_missing(self, embeddings_body: CreateEmbeddingsBody): + # Arrange + embeddings_body = CreateEmbeddingsBody(model=embeddings_body.model, input=None, messages=None) + + # Act + result = embeddings_body.get_prompts() + + # Assert + assert result == [] diff --git a/api/tests/unit/infrastructure/http/adapters/test_embeddingsadapter.py b/api/tests/unit/infrastructure/http/adapters/test_embeddingsadapter.py index 72a73a02b..2c2c83350 100644 --- a/api/tests/unit/infrastructure/http/adapters/test_embeddingsadapter.py +++ b/api/tests/unit/infrastructure/http/adapters/test_embeddingsadapter.py @@ -181,6 +181,24 @@ def test_format_request_replace_model_by_provider_model_name(self, adapter, prov assert "model" in result.body assert result.body["model"] == provider_model_name + @pytest.mark.parametrize( + argnames=("adapter"), + argvalues=["tei_embeddings_adapter", "vllm_embeddings_adapter"], + indirect=["adapter"], + ) + def test_format_request_excludes_none_fields(self, adapter): + # Arrange + original_request = ProviderOriginalRequestFactory(embeddings=True) + original_request.body.dimensions = None + + # Act + result = adapter.format_request(original_request) + + # Assert + assert "dimensions" not in result.body + assert "model" in result.body + assert "input" in result.body + @pytest.mark.parametrize( argnames=("adapter", "method"), argvalues=[("tei_embeddings_adapter", HTTPMethod.POST), ("vllm_embeddings_adapter", HTTPMethod.POST)], diff --git a/api/tests/unit/infrastructure/http/adapters/test_rerankadapter.py b/api/tests/unit/infrastructure/http/adapters/test_rerankadapter.py index f17b777de..2c8a6b182 100644 --- a/api/tests/unit/infrastructure/http/adapters/test_rerankadapter.py +++ b/api/tests/unit/infrastructure/http/adapters/test_rerankadapter.py @@ -180,6 +180,23 @@ def test_format_request_replace_model_by_provider_model_name(self, adapter, prov assert "model" in result.body assert result.body["model"] == provider_model_name + @pytest.mark.parametrize( + argnames=("adapter"), + argvalues=["tei_rerank_adapter", "vllm_rerank_adapter"], + indirect=["adapter"], + ) + def test_format_request_excludes_none_fields(self, adapter): + # Arrange + original_request = ProviderOriginalRequestFactory(rerank=True) + original_request.body.top_n = None + + # Act + result = adapter.format_request(original_request) + + # Assert + assert "top_n" not in result.body + assert "query" in result.body + @pytest.mark.parametrize( argnames=("adapter", "method"), argvalues=[("tei_rerank_adapter", HTTPMethod.POST), ("vllm_rerank_adapter", HTTPMethod.POST)], diff --git a/api/tests/unit/use_case/embeddings/test_createembeddingsusecase.py b/api/tests/unit/use_case/embeddings/test_createembeddingsusecase.py index c78e6423f..51a0becb3 100644 --- a/api/tests/unit/use_case/embeddings/test_createembeddingsusecase.py +++ b/api/tests/unit/use_case/embeddings/test_createembeddingsusecase.py @@ -749,3 +749,44 @@ async def test_should_enrich_when_non_admin_user_and_flow_succeeds( total_tokens=15, cost=result.data.usage.cost, ) + + @pytest.mark.asyncio + async def test_should_forward_extra_fields_to_provider_adapter_without_raising( + self, + use_case, + provider_adapter_builder, + request_context, + admin_user, + mock_successful_embeddings_flow, + ): + # Arrange + messages = [ + { + "role": "user", + "content": [ + {"type": "image_url", "image_url": {"url": "https://example.com/image.jpg"}}, + {"type": "text", "text": "Represent the given image."}, + ], + } + ] + request_context.set(RequestContext(user=admin_user)) + command = CreateEmbeddingsCommand( + input=None, + model="embeddings-router", + messages=messages, + continue_final_message=True, + add_special_tokens=True, + request_context=request_context, + ) + + # Act + result = await use_case.execute(command=command) + + # Assert + assert isinstance(result, CreateEmbeddingsUseCaseSuccess) + original_request = provider_adapter_builder.build.return_value.format_request.call_args.kwargs["original_request"] + body = original_request.body.model_dump() + assert body["messages"] == messages + assert body["continue_final_message"] is True + assert body["add_special_tokens"] is True + assert body["input"] is None diff --git a/api/tests/unit/use_case/reranks/test_creatererankusecase.py b/api/tests/unit/use_case/reranks/test_creatererankusecase.py index e9afc8335..a40864127 100644 --- a/api/tests/unit/use_case/reranks/test_creatererankusecase.py +++ b/api/tests/unit/use_case/reranks/test_creatererankusecase.py @@ -748,3 +748,40 @@ async def test_should_enrich_when_non_admin_user_and_flow_succeeds( total_tokens=15, cost=result.data.usage.cost, ) + + @pytest.mark.asyncio + async def test_should_forward_extra_fields_to_provider_adapter_without_raising( + self, + use_case, + provider_adapter_builder, + request_context, + admin_user, + mock_successful_rerank_flow, + ): + # Arrange + request_context.set(RequestContext(user=admin_user)) + command = CreateRerankCommand( + query="query", + documents=["doc a", "doc b"], + model="rerank-router", + top_n=2, + truncate=True, + return_text=True, + raw_scores=False, + truncation_direction="left", + request_context=request_context, + ) + + # Act + result = await use_case.execute(command=command) + + # Assert + assert isinstance(result, CreateRerankUseCaseSuccess) + original_request = provider_adapter_builder.build.return_value.format_request.call_args.kwargs["original_request"] + body = original_request.body.model_dump() + assert body["truncate"] is True + assert body["return_text"] is True + assert body["raw_scores"] is False + assert body["truncation_direction"] == "left" + assert body["query"] == "query" + assert body["documents"] == ["doc a", "doc b"] diff --git a/api/use_cases/embeddings/_createembeddingsusecase.py b/api/use_cases/embeddings/_createembeddingsusecase.py index 2b4bdb75f..87545fa6f 100644 --- a/api/use_cases/embeddings/_createembeddingsusecase.py +++ b/api/use_cases/embeddings/_createembeddingsusecase.py @@ -24,7 +24,7 @@ class CreateEmbeddingsCommand(CreateEmbeddingsBody): - model_config = ConfigDict(arbitrary_types_allowed=True) + model_config = ConfigDict(arbitrary_types_allowed=True, extra="allow") request_context: ContextVar[RequestContext] @@ -112,12 +112,7 @@ async def execute(self, command: CreateEmbeddingsCommand) -> CreateEmbeddingsUse adapter = self.provider_adapter_builder.build(endpoint=EndpointRoute.EMBEDDINGS, provider=provider) original_request = ProviderOriginalRequest( endpoint=EndpointRoute.EMBEDDINGS, - body=CreateEmbeddingsBody( - input=command.input, - model=command.model, - dimensions=command.dimensions, - encoding_format=command.encoding_format, - ), + body=CreateEmbeddingsBody.model_validate(command.model_dump(exclude={"request_context"})), ) prompt_tokens = self.model_tokenizer.compute_tokens(texts=original_request.body.get_prompts()) diff --git a/api/use_cases/reranks/_creatererankusecase.py b/api/use_cases/reranks/_creatererankusecase.py index 28e7158be..c6d1ed830 100644 --- a/api/use_cases/reranks/_creatererankusecase.py +++ b/api/use_cases/reranks/_creatererankusecase.py @@ -24,7 +24,7 @@ class CreateRerankCommand(CreateRerankBody): - model_config = ConfigDict(arbitrary_types_allowed=True) + model_config = ConfigDict(arbitrary_types_allowed=True, extra="allow") request_context: ContextVar[RequestContext] @@ -112,7 +112,7 @@ async def execute(self, command: CreateRerankCommand) -> CreateRerankUseCaseResu adapter = self.provider_adapter_builder.build(endpoint=EndpointRoute.RERANK, provider=provider) original_request = ProviderOriginalRequest( endpoint=EndpointRoute.RERANK, - body=CreateRerankBody(query=command.query, documents=command.documents, model=command.model, top_n=command.top_n), + body=CreateRerankBody.model_validate(command.model_dump(exclude={"request_context"})), ) prompt_tokens = self.model_tokenizer.compute_tokens(texts=original_request.body.get_prompts())