Skip to content

Commit dad44f5

Browse files
[pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
1 parent 2de1933 commit dad44f5

3 files changed

Lines changed: 7 additions & 4 deletions

File tree

src/litserve/specs/openai_embedding.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -146,7 +146,8 @@ def setup(self, server: "LitServer"):
146146

147147
def _check_lit_api(self, api):
148148
from litserve import LitAPI
149-
if isinstance(api.spec,OpenAIEmbeddingSpec):
149+
150+
if isinstance(api.spec, OpenAIEmbeddingSpec):
150151
if inspect.isgeneratorfunction(api.predict):
151152
raise ValueError(
152153
"You are using yield in your predict method, which is used for streaming.",

src/litserve/test_examples/openai_embedding_spec_example.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -16,6 +16,7 @@ def predict(self, x) -> List[List[float]]:
1616
def encode_response(self, output) -> dict:
1717
return {"embeddings": output}
1818

19+
1920
class TestOpenAPI(LitAPI):
2021
def setup(self, device):
2122
self.model = None
@@ -24,7 +25,7 @@ async def predict(self, x) -> List[List[float]]:
2425
n = len(x) if isinstance(x, list) else 1
2526
yield np.random.rand(n, 768).tolist()
2627

27-
async def encode_response(self, output) -> dict:
28+
async def encode_response(self, output) -> dict:
2829
yield {"embeddings": output}
2930

3031

tests/unit/test_openai_embedding.py

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -26,12 +26,12 @@
2626
from litserve.specs.openai_embedding import OpenAIEmbeddingSpec
2727
from litserve.test_examples.openai_embedding_spec_example import (
2828
TestEmbedAPI,
29-
TestOpenAPI,
3029
TestEmbedAPIWithMissingEmbeddings,
3130
TestEmbedAPIWithNonDictOutput,
3231
TestEmbedAPIWithUsage,
3332
TestEmbedAPIWithYieldEncodeResponse,
3433
TestEmbedAPIWithYieldPredict,
34+
TestOpenAPI,
3535
)
3636
from litserve.utils import wrap_litserve_start
3737

@@ -52,11 +52,12 @@ async def test_openai_embedding_spec_with_single_input(openai_embedding_request_
5252
assert len(resp.json()["data"]) == 1, "Length of data should be 1"
5353
assert len(resp.json()["data"][0]["embedding"]) == 768, "Embedding length should be 768"
5454

55+
5556
@pytest.mark.asyncio
5657
async def test_openai_embedding_spec_with_multi_endpoint(openai_embedding_request_data):
5758
spec_openai = OpenAISpec()
5859
spec_embedding = OpenAIEmbeddingSpec()
59-
server = ls.LitServer([TestOpenAPI(spec=spec_openai,enable_async=True),TestEmbedAPI(spec=spec_embedding)])
60+
server = ls.LitServer([TestOpenAPI(spec=spec_openai, enable_async=True), TestEmbedAPI(spec=spec_embedding)])
6061
with wrap_litserve_start(server) as server:
6162
async with LifespanManager(server.app) as manager, AsyncClient(
6263
transport=ASGITransport(app=manager.app), base_url="http://test"

0 commit comments

Comments
 (0)