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:
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)
@@ -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()