fix(avatar): normalize embedding API endpoint

This commit is contained in:
stefanfeng
2026-08-28 16:59:22 +08:00
parent 7884430b3d
commit 6794e88d53
3 changed files with 29 additions and 2 deletions
@@ -48,9 +48,11 @@ 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({
@@ -61,7 +63,7 @@ class RemoteEmbeddingTests(unittest.TestCase):
})
with patch.dict(os.environ, {
"EMBEDDING_API_URL": "https://embedding.example/v1/embeddings",
"EMBEDDING_API_URL": "https://embedding.example/v1",
"EMBEDDING_API_KEY": "test-key",
"EMBEDDING_MODEL": "text-embedding-v4",
"EMBEDDING_BATCH_SIZE": "10",
@@ -69,8 +71,18 @@ class RemoteEmbeddingTests(unittest.TestCase):
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()