diff --git a/digital-avatar-app/backend/embeddings.py b/digital-avatar-app/backend/embeddings.py index 6c54232..0e20019 100644 --- a/digital-avatar-app/backend/embeddings.py +++ b/digital-avatar-app/backend/embeddings.py @@ -51,22 +51,32 @@ def embed(texts): if api_url: api_key = os.getenv("EMBEDDING_API_KEY", "") model = os.getenv("EMBEDDING_MODEL", "text-embedding-3-small") - payload = json.dumps({"input": texts, "model": model}).encode("utf-8") - req = urllib.request.Request( - api_url, - data=payload, - headers={ - "Content-Type": "application/json", - "Authorization": f"Bearer {api_key}" if api_key else "", - }, - method="POST", - ) - with urllib.request.urlopen(req, timeout=30) as resp: - data = json.loads(resp.read().decode("utf-8")) - items = data["data"] - if items and "index" in items[0]: - items = sorted(items, key=lambda x: x["index"]) - return [item["embedding"] for item in items] + try: + batch_size = max(1, int(os.getenv("EMBEDDING_BATCH_SIZE", "10"))) + except ValueError: + batch_size = 10 + embeddings = [] + for start in range(0, len(texts), batch_size): + batch = texts[start:start + batch_size] + payload = json.dumps({"input": batch, "model": model}).encode("utf-8") + req = urllib.request.Request( + api_url, + data=payload, + headers={ + "Content-Type": "application/json", + "Authorization": f"Bearer {api_key}" if api_key else "", + }, + method="POST", + ) + with urllib.request.urlopen(req, timeout=30) as resp: + data = json.loads(resp.read().decode("utf-8")) + items = data["data"] + if items and "index" in items[0]: + items = sorted(items, key=lambda x: x["index"]) + if len(items) != len(batch): + raise ValueError("embedding response count does not match request") + embeddings.extend(item["embedding"] for item in items) + return embeddings return _hash_embedding(texts) diff --git a/digital-avatar-app/backend/tests/test_embeddings.py b/digital-avatar-app/backend/tests/test_embeddings.py index 7520631..eafb1d8 100644 --- a/digital-avatar-app/backend/tests/test_embeddings.py +++ b/digital-avatar-app/backend/tests/test_embeddings.py @@ -1,10 +1,26 @@ +import json import os import tempfile import unittest +from unittest.mock import patch import embeddings +class FakeResponse: + def __init__(self, payload): + self.payload = payload + + def __enter__(self): + return self + + def __exit__(self, *_): + return None + + def read(self): + return json.dumps(self.payload).encode("utf-8") + + class TextExtractionTests(unittest.TestCase): def write_text(self, suffix, content): handle = tempfile.NamedTemporaryFile(suffix=suffix, delete=False) @@ -28,5 +44,33 @@ class TextExtractionTests(unittest.TestCase): embeddings.extract_text(path, ".csv") +class RemoteEmbeddingTests(unittest.TestCase): + def test_large_input_is_split_into_provider_safe_batches(self): + texts = [f"chunk-{index}" for index in range(14)] + batch_sizes = [] + + def fake_urlopen(request, timeout): + self.assertEqual(timeout, 30) + payload = json.loads(request.data.decode("utf-8")) + batch_sizes.append(len(payload["input"])) + return FakeResponse({ + "data": [ + {"index": index, "embedding": [float(text.split("-")[1])]} + for index, text in enumerate(payload["input"]) + ] + }) + + with patch.dict(os.environ, { + "EMBEDDING_API_URL": "https://embedding.example/v1/embeddings", + "EMBEDDING_API_KEY": "test-key", + "EMBEDDING_MODEL": "text-embedding-v4", + "EMBEDDING_BATCH_SIZE": "10", + }), patch("embeddings.urllib.request.urlopen", side_effect=fake_urlopen): + result = embeddings.embed(texts) + + self.assertEqual(batch_sizes, [10, 4]) + self.assertEqual(result, [[float(index)] for index in range(14)]) + + if __name__ == "__main__": unittest.main()