acl-api / tests /sync /test_embeddings_client.py
ivykopal's picture
feat: optimize database operations with batch upserts and mark inactive functionality
ee7e79e
Raw
History Blame Contribute Delete
4.54 kB
import numpy as np
import pytest
from openai import APIStatusError
from sync.embeddings_client import EmbeddingsClient
class _FakeEmbeddingDatum:
def __init__(self, embedding, index):
self.embedding = embedding
self.index = index
class _FakeEmbeddingsResponse:
def __init__(self, data):
self.data = data
def test_embed_batch_parses_response(monkeypatch):
captured = {}
def fake_create(self, model, input):
captured["model"] = model
captured["input"] = input
return _FakeEmbeddingsResponse([
_FakeEmbeddingDatum([0.2] * 1024, 1),
_FakeEmbeddingDatum([0.1] * 1024, 0),
])
monkeypatch.setattr("openai.resources.embeddings.Embeddings.create", fake_create)
client = EmbeddingsClient(base_url="http://fake", api_key="k")
result = client.embed_batch(["a", "b"])
assert captured["input"] == ["a", "b"]
assert len(result) == 2
assert result[0].shape == (1024,)
assert result[0].dtype == np.float32
assert np.isclose(result[0][0], 0.1)
assert np.isclose(result[1][0], 0.2)
def test_embed_batch_empty_input_makes_no_requests(monkeypatch):
called = {"n": 0}
def fake_create(self, model, input):
called["n"] += 1
return _FakeEmbeddingsResponse([])
monkeypatch.setattr("openai.resources.embeddings.Embeddings.create", fake_create)
client = EmbeddingsClient(base_url="http://fake", api_key="k")
assert client.embed_batch([]) == []
assert called["n"] == 0
def test_embed_batch_chunks_large_input(monkeypatch):
captured = []
def fake_create(self, model, input):
captured.append(list(input))
# index resets to 0 within each request, as on the real API
return _FakeEmbeddingsResponse([
_FakeEmbeddingDatum([float(i)] * 1024, i) for i in range(len(input))
])
monkeypatch.setattr("openai.resources.embeddings.Embeddings.create", fake_create)
monkeypatch.setattr("sync.embeddings_client._BATCH_SIZE", 32)
client = EmbeddingsClient(base_url="http://fake", api_key="k")
texts = [f"t{i}" for i in range(80)]
result = client.embed_batch(texts)
assert len(captured) == 3 # 32 + 32 + 16
assert captured[0] == texts[:32]
assert captured[1] == texts[32:64]
assert captured[2] == texts[64:]
assert len(result) == 80
# order preserved within and across chunks (each chunk's vectors come back 0..n-1)
assert [round(v[0]) for v in result] == list(range(32)) + list(range(32)) + list(range(16))
def test_embed_batch_retries_on_timeout_then_succeeds(monkeypatch):
import httpx
from openai import APITimeoutError
calls = {"n": 0}
def fake_create(self, model, input):
calls["n"] += 1
if calls["n"] == 1:
raise APITimeoutError(request=httpx.Request("POST", "http://fake/embeddings"))
return _FakeEmbeddingsResponse([_FakeEmbeddingDatum([0.5] * 1024, 0)])
monkeypatch.setattr("openai.resources.embeddings.Embeddings.create", fake_create)
monkeypatch.setattr("sync.embeddings_client.time.sleep", lambda _: None)
client = EmbeddingsClient(base_url="http://fake", api_key="k")
result = client.embed_batch(["a"])
assert calls["n"] == 2
assert np.isclose(result[0][0], 0.5)
def test_embed_batch_raises_after_exhausted_retries(monkeypatch):
import httpx
from openai import APITimeoutError
calls = {"n": 0}
def fake_create(self, model, input):
calls["n"] += 1
raise APITimeoutError(request=httpx.Request("POST", "http://fake/embeddings"))
monkeypatch.setattr("openai.resources.embeddings.Embeddings.create", fake_create)
monkeypatch.setattr("sync.embeddings_client.time.sleep", lambda _: None)
client = EmbeddingsClient(base_url="http://fake", api_key="k")
with pytest.raises(RuntimeError):
client.embed_batch(["a"])
assert calls["n"] == 3 # _MAX_RETRIES attempts
def test_embed_batch_raises_runtime_error_on_api_error(monkeypatch):
import httpx
def fake_create(self, model, input):
request = httpx.Request("POST", "http://fake/embeddings")
response = httpx.Response(500, request=request, json={"error": "boom"})
raise APIStatusError("boom", response=response, body={"error": "boom"})
monkeypatch.setattr("openai.resources.embeddings.Embeddings.create", fake_create)
client = EmbeddingsClient(base_url="http://fake", api_key="k")
with pytest.raises(RuntimeError):
client.embed_batch(["a"])