89 lines
3.0 KiB
Python
89 lines
3.0 KiB
Python
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)
|
|
handle.close()
|
|
self.addCleanup(lambda: os.path.exists(handle.name) and os.unlink(handle.name))
|
|
with open(handle.name, "w", encoding="utf-8") as stream:
|
|
stream.write(content)
|
|
return handle.name
|
|
|
|
def test_extracts_utf8_markdown(self):
|
|
path = self.write_text(".md", "# 退款规则\n\n七日内可申请退款。")
|
|
self.assertEqual(embeddings.extract_text(path, ".md"), "# 退款规则\n\n七日内可申请退款。")
|
|
|
|
def test_extracts_utf8_text(self):
|
|
path = self.write_text(".txt", "客服热线:400-123-4567")
|
|
self.assertEqual(embeddings.extract_text(path, ".txt"), "客服热线:400-123-4567")
|
|
|
|
def test_rejects_unsupported_extension(self):
|
|
path = self.write_text(".csv", "not supported")
|
|
with self.assertRaises(ValueError):
|
|
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 = []
|
|
requested_urls = []
|
|
|
|
def fake_urlopen(request, timeout):
|
|
self.assertEqual(timeout, 30)
|
|
requested_urls.append(request.full_url)
|
|
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",
|
|
"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(requested_urls, [
|
|
"https://embedding.example/v1/embeddings",
|
|
"https://embedding.example/v1/embeddings",
|
|
])
|
|
self.assertEqual(result, [[float(index)] for index in range(14)])
|
|
|
|
def test_full_embedding_endpoint_is_not_modified(self):
|
|
self.assertEqual(
|
|
embeddings._embedding_endpoint("https://embedding.example/v1/embeddings/"),
|
|
"https://embedding.example/v1/embeddings",
|
|
)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|