diff --git a/src/litserve/specs/openai_embedding.py b/src/litserve/specs/openai_embedding.py index a6640ab4d..b5f8a2e86 100644 --- a/src/litserve/specs/openai_embedding.py +++ b/src/litserve/specs/openai_embedding.py @@ -12,8 +12,10 @@ # See the License for the specific language governing permissions and # limitations under the License. import asyncio +import base64 import inspect import logging +import struct import sys import time import uuid @@ -59,7 +61,7 @@ def ensure_list(self): class Embedding(BaseModel): index: int - embedding: list[float] + embedding: Union[list[float], str] object: Literal["embedding"] = "embedding" @@ -287,6 +289,11 @@ async def embeddings_endpoint(self, request: EmbeddingRequest) -> EmbeddingRespo self._validate_response(response) data: list[Embedding] = self._handle_embedding_response(response["embeddings"], num_items) + if request.encoding_format == "base64": + for item in data: + # OpenAI clients decode base64 embeddings as little-endian float32 vectors. + packed = struct.pack(f"<{len(item.embedding)}f", *item.embedding) + item.embedding = base64.b64encode(packed).decode("ascii") usage = UsageInfo(**response) diff --git a/tests/e2e/test_e2e.py b/tests/e2e/test_e2e.py index cd47b7432..7215bce5a 100644 --- a/tests/e2e/test_e2e.py +++ b/tests/e2e/test_e2e.py @@ -11,6 +11,7 @@ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # See the License for the specific language governing permissions and # limitations under the License. +import base64 import json import os import subprocess @@ -392,6 +393,13 @@ def test_openai_embedding_parity(): for data in response.data: assert len(data.embedding) == 768, f"Expected 768 dimensions but got {len(data.embedding)}" + # the SDK returns the raw payload when the caller explicitly asks for base64 + response = client.embeddings.create(model="lit", input=input_text, encoding_format="base64") + embedding = response.data[0].embedding + assert isinstance(embedding, str), "Expected a base64 string but got something else" + # decoding gives raw bytes, so the length pins both the dimensions and the 4-byte float32 width + assert len(base64.b64decode(embedding)) == 768 * 4, "Expected 768 little-endian float32 values" + @e2e_from_file("tests/e2e/default_async_streaming.py") def test_e2e_default_async_streaming(): diff --git a/tests/unit/test_openai_embedding.py b/tests/unit/test_openai_embedding.py index 6a45fa633..17f2a052f 100644 --- a/tests/unit/test_openai_embedding.py +++ b/tests/unit/test_openai_embedding.py @@ -13,6 +13,7 @@ # limitations under the License. import asyncio +import base64 import copy import time @@ -209,3 +210,23 @@ async def test_batching_with_client_side_batching(openai_embedding_request_data_ == "The OpenAIEmbedding spec does not support dynamic batching when client-side batching is used. " "To resolve this, either set `max_batch_size=1` or send a single input from the client." ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("request_data", ["openai_embedding_request_data", "openai_embedding_request_data_array"]) +async def test_openai_embedding_spec_with_base64_encoding(request_data, request): + request_data = {**request.getfixturevalue(request_data), "encoding_format": "base64"} + server = ls.LitServer(TestEmbedAPI(spec=OpenAIEmbeddingSpec())) + + with wrap_litserve_start(server) as server: + async with ( + LifespanManager(server.app) as manager, + AsyncClient(transport=ASGITransport(app=manager.app), base_url="http://test") as ac, + ): + resp = await ac.post("/v1/embeddings", json=request_data, timeout=10) + assert resp.status_code == 200, "Status code should be 200" + for item in resp.json()["data"]: + embedding = item["embedding"] + assert isinstance(embedding, str), "Embedding should be a base64 string" + # OpenAI clients decode base64 embeddings as little-endian float32 vectors + assert len(np.frombuffer(base64.b64decode(embedding), dtype="