import java.util.*; /** * Unit test simulating LangChain langchain-0001: * MultiVectorRetriever._get_relevant_documents() uses a list for ID dedup: * ids = [] * for d in sub_docs: * if id_key in d.metadata and d.metadata[id_key] not in ids: * ids.append(d.metadata[id_key]) * * "id not in ids" is O(|ids|) per document, making the full dedup O(k^2) * where k = number of sub-docs from vectorstore search. * Fix: use a set for O(1) membership, keeping ids list for order preservation. */ public class LangChainMultiVectorDedupTest { // Defective: O(k^2) list-based dedup static List dedupList(List subDocIds) { List ids = new ArrayList<>(); for (String id : subDocIds) { if (!ids.contains(id)) { // O(|ids|) scan ids.add(id); } } return ids; } // Fixed: O(k) set-assisted dedup, preserving order static List dedupSet(List subDocIds) { Set seen = new LinkedHashSet<>(); List ids = new ArrayList<>(); for (String id : subDocIds) { if (seen.add(id)) { // O(1) amortized ids.add(id); } } return ids; } // Measure operation count for list-based dedup static long countListOps(List subDocIds) { List ids = new ArrayList<>(); long ops = 0; for (String id : subDocIds) { ops += ids.size(); // each contains() scans whole list if (!ids.contains(id)) { ids.add(id); } } return ops; } // Measure operation count for set-based dedup static long countSetOps(List subDocIds) { Set seen = new HashSet<>(); long ops = 0; for (String id : subDocIds) { ops += 1; // O(1) hash lookup seen.add(id); } return ops; } // Simulate sub_docs where each doc has an id pointing to a parent document // k sub-docs may map to fewer unique parent ids (many-to-one) static List makeSubDocIds(int k, int uniqueParents) { Random rng = new Random(42); List ids = new ArrayList<>(k); for (int i = 0; i < k; i++) { ids.add("parent-" + (rng.nextInt(uniqueParents))); } return ids; } static void append(List list, String val) { list.add(val); } // Override dedupList to use add not append static List dedupListFixed(List subDocIds) { List ids = new ArrayList<>(); for (String id : subDocIds) { if (!ids.contains(id)) { ids.add(id); } } return ids; } public static void main(String[] args) { System.out.println("langchain-0001: MultiVectorRetriever ID dedup O(k^2) -> O(k)"); System.out.println("=".repeat(60)); // Test correctness List input = Arrays.asList( "p1", "p2", "p1", "p3", "p2", "p4", "p1" ); List expected = Arrays.asList("p1", "p2", "p3", "p4"); List listResult = dedupListFixed(input); List setResult = dedupSet(input); if (!listResult.equals(expected)) { throw new AssertionError("List dedup wrong: " + listResult); } if (!setResult.equals(expected)) { throw new AssertionError("Set dedup wrong: " + setResult); } System.out.println("PASS: both produce correct ordered dedup"); // Benchmark comparison at various k values System.out.printf("%n%-8s %12s %12s %10s%n", "k", "list_ops", "set_ops", "ratio"); System.out.println("-".repeat(44)); int[] kValues = {10, 50, 100, 500, 1000}; for (int k : kValues) { // Worst case: all unique ids (no duplicates) maximizes list scan ops List allUnique = new ArrayList<>(k); for (int i = 0; i < k; i++) allUnique.add("p" + i); long listOps = countListOps(allUnique); long setOps = countSetOps(allUnique); double ratio = (double) listOps / setOps; System.out.printf("%-8d %12d %12d %9.1fx%n", k, listOps, setOps, ratio); if (k >= 100 && ratio < 40) { throw new AssertionError("Expected >40x ratio at k=" + k + ", got " + ratio); } } // Realistic scenario: k=100 sub-docs mapping to 20 unique parents int k = 100, parents = 20; List realistic = makeSubDocIds(k, parents); long listOpsR = countListOps(realistic); long setOpsR = countSetOps(realistic); double ratioR = (double) listOpsR / setOpsR; System.out.printf("%nRealistic k=%d, %d parents: list=%d set=%d ratio=%.1fx%n", k, parents, listOpsR, setOpsR, ratioR); if (ratioR < 5) { throw new AssertionError("Expected >5x ratio in realistic case, got " + ratioR); } System.out.println("\nAll assertions PASS"); System.out.println("Fix: replace 'ids = []; id not in ids' with set-assisted dedup"); System.out.println(" seen_ids = set()"); System.out.println(" if id not in seen_ids: seen_ids.add(id); ids.append(id)"); } }