Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion .github/badges/coverage.json
Original file line number Diff line number Diff line change
@@ -1 +1 @@
{"schemaVersion": 1, "label": "coverage", "message": "59.82%", "color": "red"}
{"schemaVersion": 1, "label": "coverage", "message": "59.83%", "color": "red"}
23 changes: 21 additions & 2 deletions api/domain/embeddings/entities.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,31 +4,50 @@
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

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):
Expand Down
8 changes: 1 addition & 7 deletions api/infrastructure/fastapi/endpoints/embeddings.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
8 changes: 1 addition & 7 deletions api/infrastructure/fastapi/endpoints/rerank.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
10 changes: 9 additions & 1 deletion api/infrastructure/fastapi/schemas/embeddings.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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

Expand Down
2 changes: 1 addition & 1 deletion api/infrastructure/http/adapters/_httpprovideradapter.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 {},
)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
32 changes: 32 additions & 0 deletions api/tests/integration/endpoints/test_embeddings.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,4 @@
import json
from unittest.mock import AsyncMock, MagicMock

from httpx import AsyncClient
Expand Down Expand Up @@ -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",
[
Expand Down
32 changes: 32 additions & 0 deletions api/tests/integration/endpoints/test_rerank.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,4 @@
import json
from unittest.mock import AsyncMock, MagicMock

from httpx import AsyncClient
Expand Down Expand Up @@ -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",
[
Expand Down
8 changes: 4 additions & 4 deletions api/tests/integration/endpoints/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
55 changes: 55 additions & 0 deletions api/tests/unit/domain/embeddings/test_embeddingsentities.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 == []
Original file line number Diff line number Diff line change
Expand Up @@ -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)],
Expand Down
17 changes: 17 additions & 0 deletions api/tests/unit/infrastructure/http/adapters/test_rerankadapter.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)],
Expand Down
Loading