test_embeddings.cpp 12 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295
  1. #include <gtest/gtest.h>
  2. #include "embeddings.hpp"
  3. #include "embedding_cache.hpp"
  4. #include "errors.hpp"
  5. #include <httplib.h>
  6. #include <atomic>
  7. #include <chrono>
  8. #include <thread>
  9. #include <vector>
  10. using namespace svapi;
  11. TEST(Embeddings, ParsesWellFormedResponse) {
  12. auto v = parseEmbeddingResponse(nlohmann::json::parse(R"({"data":[{"embedding":[0.5,1.0,1.5]}]})"));
  13. ASSERT_EQ(v.size(), 3u); EXPECT_FLOAT_EQ(v[0], 0.5f); EXPECT_FLOAT_EQ(v[2], 1.5f);
  14. }
  15. TEST(Embeddings, MalformedThrows) {
  16. EXPECT_THROW(parseEmbeddingResponse(nlohmann::json::parse(R"({"data":[]})")), ApiError);
  17. EXPECT_THROW(parseEmbeddingResponse(nlohmann::json::object()), ApiError);
  18. }
  19. TEST(Embeddings, SplitBaseHandlesPathPrefix) {
  20. EXPECT_EQ(splitEmbeddingBase("https://api.openai.com").origin, "https://api.openai.com");
  21. EXPECT_EQ(splitEmbeddingBase("https://api.openai.com").path, "");
  22. EXPECT_EQ(splitEmbeddingBase("https://api.openai.com/").path, "");
  23. auto orr = splitEmbeddingBase("https://openrouter.ai/api");
  24. EXPECT_EQ(orr.origin, "https://openrouter.ai");
  25. EXPECT_EQ(orr.path, "/api");
  26. EXPECT_EQ(splitEmbeddingBase("https://openrouter.ai/api/").path, "/api");
  27. EXPECT_EQ(splitEmbeddingBase("http://h:8080/a/b").origin, "http://h:8080");
  28. EXPECT_EQ(splitEmbeddingBase("http://h:8080/a/b").path, "/a/b");
  29. }
  30. TEST(Embeddings, ClientHitsMockServer) {
  31. httplib::Server mock;
  32. mock.Post("/v1/embeddings", [](const httplib::Request& req, httplib::Response& res) {
  33. auto j = nlohmann::json::parse(req.body);
  34. EXPECT_EQ(j["model"], "m"); EXPECT_EQ(j["input"], "hello");
  35. res.set_content(R"({"data":[{"embedding":[1,2,3,4]}]})", "application/json");
  36. });
  37. int port = mock.bind_to_any_port("127.0.0.1");
  38. std::thread t([&]{ mock.listen_after_bind(); });
  39. EmbeddingClient cli("http://127.0.0.1:" + std::to_string(port), "k");
  40. auto v = cli.embed("m", "hello");
  41. EXPECT_EQ(v.size(), 4u); EXPECT_FLOAT_EQ(v[3], 4.0f);
  42. mock.stop(); t.join();
  43. }
  44. TEST(Embeddings, ClientHonorsBasePathPrefix) {
  45. // Provider whose embeddings live under a path prefix (e.g. OpenRouter /api/v1/embeddings)
  46. httplib::Server mock;
  47. mock.Post("/api/v1/embeddings", [](const httplib::Request&, httplib::Response& res) {
  48. res.set_content(R"({"data":[{"embedding":[9,9]}]})", "application/json");
  49. });
  50. int port = mock.bind_to_any_port("127.0.0.1");
  51. std::thread t([&]{ mock.listen_after_bind(); });
  52. EmbeddingClient cli("http://127.0.0.1:" + std::to_string(port) + "/api", "k");
  53. auto v = cli.embed("qwen/qwen3-embedding-8b", "hello");
  54. EXPECT_EQ(v.size(), 2u); EXPECT_FLOAT_EQ(v[0], 9.0f);
  55. mock.stop(); t.join();
  56. }
  57. TEST(Embeddings, PoolRunsRequestsConcurrently) {
  58. // A serialized impl (single Client + mutex) would pin maxInFlight at 1.
  59. // A real pool of N clients lets up to N requests overlap on independent sockets.
  60. std::atomic<int> inFlight{0};
  61. std::atomic<int> maxInFlight{0};
  62. httplib::Server mock;
  63. mock.Post("/v1/embeddings", [&](const httplib::Request&, httplib::Response& res) {
  64. int now = ++inFlight;
  65. int prev = maxInFlight.load();
  66. while (now > prev && !maxInFlight.compare_exchange_weak(prev, now)) {}
  67. std::this_thread::sleep_for(std::chrono::milliseconds(180));
  68. --inFlight;
  69. res.set_content(R"({"data":[{"embedding":[1,2,3,4]}]})", "application/json");
  70. });
  71. int port = mock.bind_to_any_port("127.0.0.1");
  72. mock.set_keep_alive_max_count(1000);
  73. std::thread srv([&]{ mock.listen_after_bind(); });
  74. while (!mock.is_running()) std::this_thread::sleep_for(std::chrono::milliseconds(5));
  75. const std::string origin = "http://127.0.0.1:" + std::to_string(port);
  76. constexpr int K = 4;
  77. std::vector<std::thread> workers;
  78. std::atomic<int> ok{0};
  79. for (int i = 0; i < K; ++i) {
  80. workers.emplace_back([&]{
  81. EmbeddingClient cli(origin, "", 10, 60, /*poolSize=*/4);
  82. auto v = cli.embed("m", "t");
  83. if (v.size() == 4u && v[3] == 4.0f) ++ok;
  84. });
  85. }
  86. for (auto& w : workers) w.join();
  87. mock.stop(); srv.join();
  88. EXPECT_EQ(ok.load(), K); // all requests succeeded
  89. EXPECT_GE(maxInFlight.load(), 2); // proves overlap — pool did NOT serialize
  90. }
  91. TEST(Embeddings, PoolSizeOneSerializes) {
  92. // poolSize=1 means a single keep-alive client → requests cannot overlap.
  93. std::atomic<int> inFlight{0};
  94. std::atomic<int> maxInFlight{0};
  95. httplib::Server mock;
  96. mock.Post("/v1/embeddings", [&](const httplib::Request&, httplib::Response& res) {
  97. int now = ++inFlight;
  98. int prev = maxInFlight.load();
  99. while (now > prev && !maxInFlight.compare_exchange_weak(prev, now)) {}
  100. std::this_thread::sleep_for(std::chrono::milliseconds(120));
  101. --inFlight;
  102. res.set_content(R"({"data":[{"embedding":[1,2,3,4]}]})", "application/json");
  103. });
  104. int port = mock.bind_to_any_port("127.0.0.1");
  105. mock.set_keep_alive_max_count(1000);
  106. std::thread srv([&]{ mock.listen_after_bind(); });
  107. while (!mock.is_running()) std::this_thread::sleep_for(std::chrono::milliseconds(5));
  108. const std::string origin = "http://127.0.0.1:" + std::to_string(port);
  109. constexpr int K = 4;
  110. std::vector<std::thread> workers;
  111. std::atomic<int> ok{0};
  112. for (int i = 0; i < K; ++i) {
  113. workers.emplace_back([&]{
  114. EmbeddingClient cli(origin, "", 10, 60, /*poolSize=*/1);
  115. auto v = cli.embed("m", "t");
  116. if (v.size() == 4u) ++ok;
  117. });
  118. }
  119. for (auto& w : workers) w.join();
  120. mock.stop(); srv.join();
  121. EXPECT_EQ(ok.load(), K);
  122. EXPECT_EQ(maxInFlight.load(), 1); // serialized by design (single client)
  123. }
  124. // ---- EmbeddingCache tests -----------------------------------------------
  125. TEST(EmbeddingCache, ReturnsCachedVectorOnExactKey) {
  126. EmbeddingCache cache(16, 1024*1024, 0);
  127. EmbeddingCache::Key k{"https://api.openai.com", "text-embedding-3-small", "hello world"};
  128. std::vector<float> vec = {0.1f, 0.2f, 0.3f};
  129. cache.put(k, vec);
  130. auto result = cache.get(k);
  131. ASSERT_TRUE(result.has_value());
  132. ASSERT_EQ(result->size(), 3u);
  133. EXPECT_FLOAT_EQ((*result)[0], 0.1f);
  134. EXPECT_FLOAT_EQ((*result)[1], 0.2f);
  135. EXPECT_FLOAT_EQ((*result)[2], 0.3f);
  136. }
  137. TEST(EmbeddingCache, MissOnAbsentKey) {
  138. EmbeddingCache cache(16, 1024*1024, 0);
  139. EmbeddingCache::Key k{"https://api.openai.com", "text-embedding-3-small", "not stored"};
  140. auto result = cache.get(k);
  141. EXPECT_FALSE(result.has_value());
  142. auto s = cache.stats();
  143. EXPECT_EQ(s.misses, 1u);
  144. EXPECT_EQ(s.hits, 0u);
  145. }
  146. TEST(EmbeddingCache, EvictsLeastRecentlyUsedWhenFull) {
  147. // capacity=3: insert A,B,C; then get(A) to refresh it; insert D → B should be evicted
  148. EmbeddingCache cache(3, 1024*1024, 0);
  149. EmbeddingCache::Key kA{"base", "m", "A"};
  150. EmbeddingCache::Key kB{"base", "m", "B"};
  151. EmbeddingCache::Key kC{"base", "m", "C"};
  152. EmbeddingCache::Key kD{"base", "m", "D"};
  153. cache.put(kA, {1.0f});
  154. cache.put(kB, {2.0f});
  155. cache.put(kC, {3.0f});
  156. // Touch A so it's recently used; B becomes the LRU
  157. cache.get(kA);
  158. // Insert D — evicts B (least recently used)
  159. cache.put(kD, {4.0f});
  160. EXPECT_TRUE(cache.get(kA).has_value());
  161. EXPECT_FALSE(cache.get(kB).has_value()); // evicted
  162. EXPECT_TRUE(cache.get(kC).has_value());
  163. EXPECT_TRUE(cache.get(kD).has_value());
  164. }
  165. TEST(EmbeddingCache, DistinguishesModels) {
  166. EmbeddingCache cache(16, 1024*1024, 0);
  167. EmbeddingCache::Key k1{"https://api.openai.com", "model-A", "hello"};
  168. EmbeddingCache::Key k2{"https://api.openai.com", "model-B", "hello"};
  169. cache.put(k1, {1.0f});
  170. cache.put(k2, {2.0f});
  171. auto r1 = cache.get(k1);
  172. auto r2 = cache.get(k2);
  173. ASSERT_TRUE(r1.has_value());
  174. ASSERT_TRUE(r2.has_value());
  175. EXPECT_FLOAT_EQ((*r1)[0], 1.0f);
  176. EXPECT_FLOAT_EQ((*r2)[0], 2.0f);
  177. EXPECT_EQ(cache.stats().size, 2u);
  178. }
  179. TEST(EmbeddingCache, DistinguishesApiBase) {
  180. EmbeddingCache cache(16, 1024*1024, 0);
  181. EmbeddingCache::Key k1{"https://api.openai.com", "text-embedding-3-small", "hello"};
  182. EmbeddingCache::Key k2{"https://openrouter.ai/api", "text-embedding-3-small", "hello"};
  183. cache.put(k1, {1.0f});
  184. cache.put(k2, {2.0f});
  185. auto r1 = cache.get(k1);
  186. auto r2 = cache.get(k2);
  187. ASSERT_TRUE(r1.has_value());
  188. ASSERT_TRUE(r2.has_value());
  189. EXPECT_FLOAT_EQ((*r1)[0], 1.0f);
  190. EXPECT_FLOAT_EQ((*r2)[0], 2.0f);
  191. EXPECT_EQ(cache.stats().size, 2u);
  192. }
  193. TEST(EmbeddingCache, ExpiresAfterTtl) {
  194. EmbeddingCache cache(16, 1024*1024, /*ttlSec=*/60);
  195. int64_t fakeNow = 1000;
  196. cache.setClockForTesting([&fakeNow]{ return fakeNow; });
  197. EmbeddingCache::Key k{"base", "m", "text"};
  198. cache.put(k, {5.0f});
  199. // Before expiry — hit
  200. EXPECT_TRUE(cache.get(k).has_value());
  201. // Advance clock past TTL
  202. fakeNow = 1061;
  203. auto result = cache.get(k);
  204. EXPECT_FALSE(result.has_value()); // expired
  205. auto s = cache.stats();
  206. EXPECT_GE(s.misses, 1u); // at least the expired-get counted as miss
  207. }
  208. TEST(EmbeddingCache, EvictsToStayUnderByteBudget) {
  209. // Each vector of 4 floats = 16 bytes. Budget = 40 bytes → max 2 full entries.
  210. EmbeddingCache cache(/*capacity=*/100, /*maxBytes=*/40, /*ttlSec=*/0);
  211. EmbeddingCache::Key kA{"b", "m", "A"};
  212. EmbeddingCache::Key kB{"b", "m", "B"};
  213. EmbeddingCache::Key kC{"b", "m", "C"};
  214. cache.put(kA, {1.0f, 2.0f, 3.0f, 4.0f}); // 16 bytes
  215. cache.put(kB, {5.0f, 6.0f, 7.0f, 8.0f}); // 16 bytes; total=32
  216. // inserting C (16 bytes) would push total to 48 > 40 → evict LRU (A)
  217. cache.put(kC, {9.0f, 10.0f, 11.0f, 12.0f});
  218. EXPECT_LE(cache.stats().bytes, 40u);
  219. EXPECT_FALSE(cache.get(kA).has_value()); // evicted
  220. EXPECT_TRUE(cache.get(kB).has_value());
  221. EXPECT_TRUE(cache.get(kC).has_value());
  222. }
  223. TEST(EmbeddingCache, NormalizationCollapsesWhitespaceAndCaseWhenEnabled) {
  224. // With normalize=true: "Max DC" and "max dc" should map to ONE entry
  225. EmbeddingCache cacheOn(16, 1024*1024, 0, /*normalize=*/true);
  226. EmbeddingCache::Key k1{"base", "m", "Max DC"};
  227. EmbeddingCache::Key k2{"base", "m", "max dc"};
  228. cacheOn.put(k1, {1.0f});
  229. auto r = cacheOn.get(k2);
  230. ASSERT_TRUE(r.has_value());
  231. EXPECT_FLOAT_EQ((*r)[0], 1.0f);
  232. EXPECT_EQ(cacheOn.stats().size, 1u);
  233. // With normalize=false (default): they stay distinct
  234. EmbeddingCache cacheOff(16, 1024*1024, 0, /*normalize=*/false);
  235. EmbeddingCache::Key k3{"base", "m", "Max DC"};
  236. EmbeddingCache::Key k4{"base", "m", "max dc"};
  237. cacheOff.put(k3, {1.0f});
  238. auto r2 = cacheOff.get(k4);
  239. EXPECT_FALSE(r2.has_value());
  240. EXPECT_EQ(cacheOff.stats().size, 1u);
  241. }
  242. TEST(EmbeddingCache, StatsCountsHitsAndMisses) {
  243. EmbeddingCache cache(16, 1024*1024, 0);
  244. EmbeddingCache::Key k{"base", "m", "text"};
  245. // 2 misses
  246. cache.get(k);
  247. cache.get(k);
  248. // put then 3 hits
  249. cache.put(k, {1.0f});
  250. cache.get(k);
  251. cache.get(k);
  252. cache.get(k);
  253. auto s = cache.stats();
  254. EXPECT_EQ(s.misses, 2u);
  255. EXPECT_EQ(s.hits, 3u);
  256. EXPECT_EQ(s.size, 1u);
  257. EXPECT_EQ(s.capacity, 16u);
  258. EXPECT_EQ(s.bytes, sizeof(float) * 1u);
  259. }
  260. TEST(EmbeddingCache, ReconfigureEvictsToTighterBounds) {
  261. EmbeddingCache cache(16, 1024*1024, 0);
  262. EmbeddingCache::Key kA{"b", "m", "A"};
  263. EmbeddingCache::Key kB{"b", "m", "B"};
  264. EmbeddingCache::Key kC{"b", "m", "C"};
  265. cache.put(kA, {1.0f}); // A is least-recently used
  266. cache.put(kB, {2.0f});
  267. cache.put(kC, {3.0f});
  268. ASSERT_EQ(cache.stats().size, 3u);
  269. // Shrink capacity to 1 — reconfigure must evict immediately, keeping the MRU entry.
  270. cache.reconfigure(1, 1024*1024, 0, false);
  271. auto s = cache.stats();
  272. EXPECT_EQ(s.size, 1u);
  273. EXPECT_EQ(s.capacity, 1u);
  274. EXPECT_FALSE(cache.get(kA).has_value()); // evicted
  275. EXPECT_TRUE(cache.get(kC).has_value()); // most-recently used survives
  276. }