codekingpro/portable-devtools
115k
1from typing import Dict, List, Set, cast
2from chromadb.types import LogRecord, Operation, Vector
3
4
5class Batch:
6 """Used to model the set of changes as an atomic operation"""
7
8 _ids_to_records: Dict[str, LogRecord]
9 _deleted_ids: Set[str]
10 _written_ids: Set[str]
11 _upsert_add_ids: Set[str] # IDs that are being added in an upsert
12 add_count: int
13 update_count: int
14
15 def __init__(self) -> None:
16 self._ids_to_records = {}
17 self._deleted_ids = set()
18 self._written_ids = set()
19 self._upsert_add_ids = set()
20 self.add_count = 0
21 self.update_count = 0
22
23 def __len__(self) -> int:
24 """Get the number of changes in this batch"""
25 return len(self._written_ids) + len(self._deleted_ids)
26
27 def get_deleted_ids(self) -> List[str]:
28 """Get the list of deleted embeddings in this batch"""
29 return list(self._deleted_ids)
30
31 def get_written_ids(self) -> List[str]:
32 """Get the list of written embeddings in this batch"""
33 return list(self._written_ids)
34
35 def get_written_vectors(self, ids: List[str]) -> List[Vector]:
36 """Get the list of vectors to write in this batch"""
37 return [
38 cast(Vector, self._ids_to_records[id]["record"]["embedding"]) for id in ids
39 ]
40
41 def get_record(self, id: str) -> LogRecord:
42 """Get the record for a given ID"""
43 return self._ids_to_records[id]
44
45 def is_deleted(self, id: str) -> bool:
46 """Check if a given ID is deleted"""
47 return id in self._deleted_ids
48
49 @property
50 def delete_count(self) -> int:
51 return len(self._deleted_ids)
52
53 def apply(self, record: LogRecord, exists_already: bool = False) -> None:
54 """Apply an embedding record to this batch. Records passed to this method are assumed to be validated for correctness.
55 For example, a delete or update presumes the ID exists in the index. An add presumes the ID does not exist in the index.
56 The exists_already flag should be set to True if the ID does exist in the index, and False otherwise.
57 """
58 id = record["record"]["id"]
59 if record["record"]["operation"] == Operation.DELETE:
60 # If the ID was previously written, remove it from the written set
61 # And update the add/update/delete counts
62 if id in self._written_ids:
63 self._written_ids.remove(id)
64 if self._ids_to_records[id]["record"]["operation"] == Operation.ADD:
65 self.add_count -= 1
66 elif (
67 self._ids_to_records[id]["record"]["operation"] == Operation.UPDATE
68 ):
69 self.update_count -= 1
70 self._deleted_ids.add(id)
71 elif (
72 self._ids_to_records[id]["record"]["operation"] == Operation.UPSERT
73 ):
74 if id in self._upsert_add_ids:
75 self.add_count -= 1
76 self._upsert_add_ids.remove(id)
77 else:
78 self.update_count -= 1
79 self._deleted_ids.add(id)
80 elif id not in self._deleted_ids:
81 self._deleted_ids.add(id)
82
83 # Remove the record from the batch
84 if id in self._ids_to_records:
85 del self._ids_to_records[id]
86
87 else:
88 self._ids_to_records[id] = record
89 self._written_ids.add(id)
90
91 # If the ID was previously deleted, remove it from the deleted set
92 # And update the delete count
93 if id in self._deleted_ids:
94 self._deleted_ids.remove(id)
95
96 # Update the add/update counts
97 if record["record"]["operation"] == Operation.UPSERT:
98 if not exists_already:
99 self.add_count += 1
100 self._upsert_add_ids.add(id)
101 else:
102 self.update_count += 1
103 elif record["record"]["operation"] == Operation.ADD:
104 self.add_count += 1
105 elif record["record"]["operation"] == Operation.UPDATE:
106 self.update_count += 1
107 