feat: 完成数字分身多分身管理与生产 H5 接入 #3
@@ -51,7 +51,14 @@ 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")
|
||||
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,
|
||||
@@ -66,7 +73,10 @@ def embed(texts):
|
||||
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]
|
||||
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()
|
||||
|
||||
Reference in New Issue
Block a user