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 = [] progress_updates = [] 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, on_progress=lambda completed, total: progress_updates.append((completed, total)), ) 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)]) self.assertEqual(progress_updates, [(10, 14), (14, 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()