fix(avatar): batch knowledge embedding requests

This commit is contained in:
stefanfeng
2026-08-21 13:21:05 +08:00
parent e2b928273c
commit e720baa21e
2 changed files with 70 additions and 16 deletions
+26 -16
View File
@@ -51,22 +51,32 @@ def embed(texts):
if api_url: if api_url:
api_key = os.getenv("EMBEDDING_API_KEY", "") api_key = os.getenv("EMBEDDING_API_KEY", "")
model = os.getenv("EMBEDDING_MODEL", "text-embedding-3-small") model = os.getenv("EMBEDDING_MODEL", "text-embedding-3-small")
payload = json.dumps({"input": texts, "model": model}).encode("utf-8") try:
req = urllib.request.Request( batch_size = max(1, int(os.getenv("EMBEDDING_BATCH_SIZE", "10")))
api_url, except ValueError:
data=payload, batch_size = 10
headers={ embeddings = []
"Content-Type": "application/json", for start in range(0, len(texts), batch_size):
"Authorization": f"Bearer {api_key}" if api_key else "", batch = texts[start:start + batch_size]
}, payload = json.dumps({"input": batch, "model": model}).encode("utf-8")
method="POST", req = urllib.request.Request(
) api_url,
with urllib.request.urlopen(req, timeout=30) as resp: data=payload,
data = json.loads(resp.read().decode("utf-8")) headers={
items = data["data"] "Content-Type": "application/json",
if items and "index" in items[0]: "Authorization": f"Bearer {api_key}" if api_key else "",
items = sorted(items, key=lambda x: x["index"]) },
return [item["embedding"] for item in items] 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) return _hash_embedding(texts)
@@ -1,10 +1,26 @@
import json
import os import os
import tempfile import tempfile
import unittest import unittest
from unittest.mock import patch
import embeddings 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): class TextExtractionTests(unittest.TestCase):
def write_text(self, suffix, content): def write_text(self, suffix, content):
handle = tempfile.NamedTemporaryFile(suffix=suffix, delete=False) handle = tempfile.NamedTemporaryFile(suffix=suffix, delete=False)
@@ -28,5 +44,33 @@ class TextExtractionTests(unittest.TestCase):
embeddings.extract_text(path, ".csv") 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__": if __name__ == "__main__":
unittest.main() unittest.main()