codekingpro/portable-devtools
115k
1from typing import Optional, Sequence, Any, Tuple, cast, Generator, Union, Dict, List
2from chromadb.segment import MetadataReader
3from chromadb.ingest import Consumer
4from chromadb.config import System
5from chromadb.types import RequestVersionContext, Segment, InclusionExclusionOperator
6from chromadb.db.impl.sqlite import SqliteDB
7from overrides import override
8from chromadb.db.base import (
9 Cursor,
10 ParameterValue,
11 get_sql,
12)
13from chromadb.telemetry.opentelemetry import (
14 OpenTelemetryClient,
15 OpenTelemetryGranularity,
16 trace_method,
17)
18from chromadb.types import (
19 Where,
20 WhereDocument,
21 MetadataEmbeddingRecord,
22 LogRecord,
23 SeqId,
24 Operation,
25 UpdateMetadata,
26 LiteralValue,
27 WhereOperator,
28)
29from uuid import UUID
30from pypika import Table, Tables
31from pypika.queries import QueryBuilder
32import pypika.functions as fn
33from pypika.terms import Criterion
34from itertools import groupby
35from functools import reduce
36import sqlite3
37
38import logging
39
40logger = logging.getLogger(__name__)
41
42
43class SqliteMetadataSegment(MetadataReader):
44 _consumer: Consumer
45 _db: SqliteDB
46 _id: UUID
47 _opentelemetry_client: OpenTelemetryClient
48 _collection_id: Optional[UUID]
49 _subscription: Optional[UUID] = None
50
51 def __init__(self, system: System, segment: Segment):
52 self._db = system.instance(SqliteDB)
53 self._consumer = system.instance(Consumer)
54 self._id = segment["id"]
55 self._opentelemetry_client = system.require(OpenTelemetryClient)
56 self._collection_id = segment["collection"]
57
58 @trace_method("SqliteMetadataSegment.start", OpenTelemetryGranularity.ALL)
59 @override
60 def start(self) -> None:
61 if self._collection_id:
62 seq_id = self.max_seqid()
63 self._subscription = self._consumer.subscribe(
64 collection_id=self._collection_id,
65 consume_fn=self._write_metadata,
66 start=seq_id,
67 )
68
69 @trace_method("SqliteMetadataSegment.stop", OpenTelemetryGranularity.ALL)
70 @override
71 def stop(self) -> None:
72 if self._subscription:
73 self._consumer.unsubscribe(self._subscription)
74
75 @trace_method("SqliteMetadataSegment.max_seqid", OpenTelemetryGranularity.ALL)
76 @override
77 def max_seqid(self) -> SeqId:
78 t = Table("max_seq_id")
79 q = (
80 self._db.querybuilder()
81 .from_(t)
82 .select(t.seq_id)
83 .where(t.segment_id == ParameterValue(self._db.uuid_to_db(self._id)))
84 )
85 sql, params = get_sql(q)
86 with self._db.tx() as cur:
87 result = cur.execute(sql, params).fetchone()
88
89 if result is None:
90 return self._consumer.min_seqid()
91 else:
92 return cast(int, result[0])
93
94 @trace_method("SqliteMetadataSegment.count", OpenTelemetryGranularity.ALL)
95 @override
96 def count(self, request_version_context: RequestVersionContext) -> int:
97 embeddings_t = Table("embeddings")
98 q = (
99 self._db.querybuilder()
100 .from_(embeddings_t)
101 .where(
102 embeddings_t.segment_id == ParameterValue(self._db.uuid_to_db(self._id))
103 )
104 .select(fn.Count(embeddings_t.id))
105 )
106 sql, params = get_sql(q)
107 with self._db.tx() as cur:
108 result = cur.execute(sql, params).fetchone()[0]
109 return cast(int, result)
110
111 @trace_method("SqliteMetadataSegment.get_metadata", OpenTelemetryGranularity.ALL)
112 @override
113 def get_metadata(
114 self,
115 request_version_context: RequestVersionContext,
116 where: Optional[Where] = None,
117 where_document: Optional[WhereDocument] = None,
118 ids: Optional[Sequence[str]] = None,
119 limit: Optional[int] = None,
120 offset: Optional[int] = None,
121 include_metadata: bool = True,
122 ) -> Sequence[MetadataEmbeddingRecord]:
123 """Query for embedding metadata."""
124 embeddings_t, metadata_t, fulltext_t = Tables(
125 "embeddings", "embedding_metadata", "embedding_fulltext_search"
126 )
127
128 limit = limit or 2**63 - 1
129 offset = offset or 0
130
131 if limit < 0:
132 raise ValueError("Limit cannot be negative")
133
134 select_clause = [
135 embeddings_t.id,
136 embeddings_t.embedding_id,
137 embeddings_t.seq_id,
138 ]
139 if include_metadata:
140 select_clause.extend(
141 [
142 metadata_t.key,
143 metadata_t.string_value,
144 metadata_t.int_value,
145 metadata_t.float_value,
146 metadata_t.bool_value,
147 ]
148 )
149
150 q = (
151 (
152 self._db.querybuilder()
153 .from_(embeddings_t)
154 .left_join(metadata_t)
155 .on(embeddings_t.id == metadata_t.id)
156 )
157 .select(*select_clause)
158 .orderby(embeddings_t.id)
159 )
160
161 # If there is a query that touches the metadata table, it uses
162 # where and where_document filters, we treat this case seperately
163 if where is not None or where_document is not None:
164 metadata_q = (
165 self._db.querybuilder()
166 .from_(embeddings_t)
167 .select(embeddings_t.id)
168 .left_join(metadata_t)
169 .on(embeddings_t.id == metadata_t.id)
170 .orderby(embeddings_t.id)
171 .where(
172 embeddings_t.segment_id
173 == ParameterValue(self._db.uuid_to_db(self._id))
174 )
175 .distinct() # These are embedding ids
176 )
177
178 if where:
179 metadata_q = metadata_q.where(
180 self._where_map_criterion(
181 metadata_q, where, metadata_t, embeddings_t
182 )
183 )
184 if where_document:
185 metadata_q = metadata_q.where(
186 self._where_doc_criterion(
187 metadata_q, where_document, metadata_t, fulltext_t, embeddings_t
188 )
189 )
190 if ids is not None:
191 metadata_q = metadata_q.where(
192 embeddings_t.embedding_id.isin(ParameterValue(ids))
193 )
194
195 metadata_q = metadata_q.limit(limit)
196 metadata_q = metadata_q.offset(offset)
197
198 q = q.where(embeddings_t.id.isin(metadata_q))
199 else:
200 # In the case where we don't use the metadata table
201 # We have to apply limit/offset to embeddings and then join
202 # with metadata
203 embeddings_q = (
204 self._db.querybuilder()
205 .from_(embeddings_t)
206 .select(embeddings_t.id)
207 .where(
208 embeddings_t.segment_id
209 == ParameterValue(self._db.uuid_to_db(self._id))
210 )
211 .orderby(embeddings_t.id)
212 .limit(limit)
213 .offset(offset)
214 )
215
216 if ids is not None:
217 embeddings_q = embeddings_q.where(
218 embeddings_t.embedding_id.isin(ParameterValue(ids))
219 )
220
221 q = q.where(embeddings_t.id.isin(embeddings_q))
222
223 with self._db.tx() as cur:
224 # Execute the query with the limit and offset already applied
225 return list(self._records(cur, q, include_metadata))
226
227 def _records(
228 self, cur: Cursor, q: QueryBuilder, include_metadata: bool
229 ) -> Generator[MetadataEmbeddingRecord, None, None]:
230 """Given a cursor and a QueryBuilder, yield a generator of records. Assumes
231 cursor returns rows in ID order."""
232
233 sql, params = get_sql(q)
234 cur.execute(sql, params)
235
236 cur_iterator = iter(cur.fetchone, None)
237 group_iterator = groupby(cur_iterator, lambda r: int(r[0]))
238
239 for _, group in group_iterator:
240 yield self._record(list(group), include_metadata)
241
242 @trace_method("SqliteMetadataSegment._record", OpenTelemetryGranularity.ALL)
243 def _record(
244 self, rows: Sequence[Tuple[Any, ...]], include_metadata: bool
245 ) -> MetadataEmbeddingRecord:
246 """Given a list of DB rows with the same ID, construct a
247 MetadataEmbeddingRecord"""
248 _, embedding_id, seq_id = rows[0][:3]
249 if not include_metadata:
250 return MetadataEmbeddingRecord(id=embedding_id, metadata=None)
251
252 metadata = {}
253 for row in rows:
254 key, string_value, int_value, float_value, bool_value = row[3:]
255 if string_value is not None:
256 metadata[key] = string_value
257 elif int_value is not None:
258 metadata[key] = int_value
259 elif float_value is not None:
260 metadata[key] = float_value
261 elif bool_value is not None:
262 if bool_value == 1:
263 metadata[key] = True
264 else:
265 metadata[key] = False
266
267 return MetadataEmbeddingRecord(
268 id=embedding_id,
269 metadata=metadata or None,
270 )
271
272 @trace_method("SqliteMetadataSegment._insert_record", OpenTelemetryGranularity.ALL)
273 def _insert_record(self, cur: Cursor, record: LogRecord, upsert: bool) -> None:
274 """Add or update a single EmbeddingRecord into the DB"""
275
276 t = Table("embeddings")
277 q = (
278 self._db.querybuilder()
279 .into(t)
280 .columns(t.segment_id, t.embedding_id, t.seq_id)
281 .where(t.segment_id == ParameterValue(self._db.uuid_to_db(self._id)))
282 .where(t.embedding_id == ParameterValue(record["record"]["id"]))
283 ).insert(
284 ParameterValue(self._db.uuid_to_db(self._id)),
285 ParameterValue(record["record"]["id"]),
286 ParameterValue(record["log_offset"]),
287 )
288 sql, params = get_sql(q)
289 sql = sql + "RETURNING id"
290 try:
291 id = cur.execute(sql, params).fetchone()[0]
292 except sqlite3.IntegrityError:
293 # Can't use INSERT OR REPLACE here because it changes the primary key.
294 if upsert:
295 return self._update_record(cur, record)
296 else:
297 logger.warning(
298 f"Insert of existing embedding ID: {record['record']['id']}"
299 )
300 # We are trying to add for a record that already exists. Fail the call.
301 # We don't throw an exception since this is in principal an async path
302 return
303
304 if record["record"]["metadata"]:
305 self._update_metadata(cur, id, record["record"]["metadata"])
306
307 @trace_method(
308 "SqliteMetadataSegment._update_metadata", OpenTelemetryGranularity.ALL
309 )
310 def _update_metadata(self, cur: Cursor, id: int, metadata: UpdateMetadata) -> None:
311 """Update the metadata for a single EmbeddingRecord"""
312 t = Table("embedding_metadata")
313 to_delete = [k for k, v in metadata.items() if v is None]
314 if to_delete:
315 q = (
316 self._db.querybuilder()
317 .from_(t)
318 .where(t.id == ParameterValue(id))
319 .where(t.key.isin(ParameterValue(to_delete)))
320 .delete()
321 )
322 sql, params = get_sql(q)
323 cur.execute(sql, params)
324
325 self._insert_metadata(cur, id, metadata)
326
327 @trace_method(
328 "SqliteMetadataSegment._insert_metadata", OpenTelemetryGranularity.ALL
329 )
330 def _insert_metadata(self, cur: Cursor, id: int, metadata: UpdateMetadata) -> None:
331 """Insert or update each metadata row for a single embedding record"""
332 t = Table("embedding_metadata")
333 q = (
334 self._db.querybuilder()
335 .into(t)
336 .columns(
337 t.id,
338 t.key,
339 t.string_value,
340 t.int_value,
341 t.float_value,
342 t.bool_value,
343 )
344 )
345 for key, value in metadata.items():
346 if isinstance(value, str):
347 q = q.insert(
348 ParameterValue(id),
349 ParameterValue(key),
350 ParameterValue(value),
351 None,
352 None,
353 None,
354 )
355 # isinstance(True, int) evaluates to True, so we need to check for bools separately
356 elif isinstance(value, bool):
357 q = q.insert(
358 ParameterValue(id),
359 ParameterValue(key),
360 None,
361 None,
362 None,
363 ParameterValue(value),
364 )
365 elif isinstance(value, int):
366 q = q.insert(
367 ParameterValue(id),
368 ParameterValue(key),
369 None,
370 ParameterValue(value),
371 None,
372 None,
373 )
374 elif isinstance(value, float):
375 q = q.insert(
376 ParameterValue(id),
377 ParameterValue(key),
378 None,
379 None,
380 ParameterValue(value),
381 None,
382 )
383
384 sql, params = get_sql(q)
385 sql = sql.replace("INSERT", "INSERT OR REPLACE")
386 if sql:
387 cur.execute(sql, params)
388
389 if "chroma:document" in metadata:
390 t = Table("embedding_fulltext_search")
391
392 def insert_into_fulltext_search() -> None:
393 q = (
394 self._db.querybuilder()
395 .into(t)
396 .columns(t.rowid, t.string_value)
397 .insert(
398 ParameterValue(id),
399 ParameterValue(metadata["chroma:document"]),
400 )
401 )
402 sql, params = get_sql(q)
403 cur.execute(sql, params)
404
405 try:
406 insert_into_fulltext_search()
407 except sqlite3.IntegrityError:
408 q = (
409 self._db.querybuilder()
410 .from_(t)
411 .where(t.rowid == ParameterValue(id))
412 .delete()
413 )
414 sql, params = get_sql(q)
415 cur.execute(sql, params)
416 insert_into_fulltext_search()
417
418 @trace_method("SqliteMetadataSegment._delete_record", OpenTelemetryGranularity.ALL)
419 def _delete_record(self, cur: Cursor, record: LogRecord) -> None:
420 """Delete a single EmbeddingRecord from the DB"""
421 t = Table("embeddings")
422 fts_t = Table("embedding_fulltext_search")
423 q = (
424 self._db.querybuilder()
425 .from_(t)
426 .where(t.segment_id == ParameterValue(self._db.uuid_to_db(self._id)))
427 .where(t.embedding_id == ParameterValue(record["record"]["id"]))
428 .delete()
429 )
430 q_fts = (
431 self._db.querybuilder()
432 .from_(fts_t)
433 .delete()
434 .where(
435 fts_t.rowid.isin(
436 self._db.querybuilder()
437 .from_(t)
438 .select(t.id)
439 .where(
440 t.segment_id == ParameterValue(self._db.uuid_to_db(self._id))
441 )
442 .where(t.embedding_id == ParameterValue(record["record"]["id"]))
443 )
444 )
445 )
446 cur.execute(*get_sql(q_fts))
447 sql, params = get_sql(q)
448 sql = sql + " RETURNING id"
449 result = cur.execute(sql, params).fetchone()
450 if result is None:
451 logger.warning(
452 f"Delete of nonexisting embedding ID: {record['record']['id']}"
453 )
454 else:
455 id = result[0]
456
457 # Manually delete metadata; cannot use cascade because
458 # that triggers on replace
459 metadata_t = Table("embedding_metadata")
460
461 q = (
462 self._db.querybuilder()
463 .from_(metadata_t)
464 .where(metadata_t.id == ParameterValue(id))
465 .delete()
466 )
467 sql, params = get_sql(q)
468 cur.execute(sql, params)
469
470 @trace_method("SqliteMetadataSegment._update_record", OpenTelemetryGranularity.ALL)
471 def _update_record(self, cur: Cursor, record: LogRecord) -> None:
472 """Update a single EmbeddingRecord in the DB"""
473 t = Table("embeddings")
474 q = (
475 self._db.querybuilder()
476 .update(t)
477 .set(t.seq_id, ParameterValue(record["log_offset"]))
478 .where(t.segment_id == ParameterValue(self._db.uuid_to_db(self._id)))
479 .where(t.embedding_id == ParameterValue(record["record"]["id"]))
480 )
481 sql, params = get_sql(q)
482 sql = sql + " RETURNING id"
483 result = cur.execute(sql, params).fetchone()
484 if result is None:
485 logger.warning(
486 f"Update of nonexisting embedding ID: {record['record']['id']}"
487 )
488 else:
489 id = result[0]
490 if record["record"]["metadata"]:
491 self._update_metadata(cur, id, record["record"]["metadata"])
492
493 @trace_method("SqliteMetadataSegment._write_metadata", OpenTelemetryGranularity.ALL)
494 def _write_metadata(self, records: Sequence[LogRecord]) -> None:
495 """Write embedding metadata to the database. Care should be taken to ensure
496 records are append-only (that is, that seq-ids should increase monotonically)"""
497 with self._db.tx() as cur:
498 for record in records:
499 if record["record"]["operation"] == Operation.ADD:
500 self._insert_record(cur, record, False)
501 elif record["record"]["operation"] == Operation.UPSERT:
502 self._insert_record(cur, record, True)
503 elif record["record"]["operation"] == Operation.DELETE:
504 self._delete_record(cur, record)
505 elif record["record"]["operation"] == Operation.UPDATE:
506 self._update_record(cur, record)
507
508 q = (
509 self._db.querybuilder()
510 .into(Table("max_seq_id"))
511 .columns("segment_id", "seq_id")
512 .insert(
513 ParameterValue(self._db.uuid_to_db(self._id)),
514 ParameterValue(record["log_offset"]),
515 )
516 )
517 sql, params = get_sql(q)
518 sql = sql.replace("INSERT", "INSERT OR REPLACE")
519 cur.execute(sql, params)
520
521 @trace_method(
522 "SqliteMetadataSegment._where_map_criterion", OpenTelemetryGranularity.ALL
523 )
524 def _where_map_criterion(
525 self, q: QueryBuilder, where: Where, metadata_t: Table, embeddings_t: Table
526 ) -> Criterion:
527 clause: List[Criterion] = []
528 for k, v in where.items():
529 if k == "$and":
530 criteria = [
531 self._where_map_criterion(q, w, metadata_t, embeddings_t)
532 for w in cast(Sequence[Where], v)
533 ]
534 clause.append(reduce(lambda x, y: x & y, criteria))
535 elif k == "$or":
536 criteria = [
537 self._where_map_criterion(q, w, metadata_t, embeddings_t)
538 for w in cast(Sequence[Where], v)
539 ]
540 clause.append(reduce(lambda x, y: x | y, criteria))
541 else:
542 expr = cast(Union[LiteralValue, Dict[WhereOperator, LiteralValue]], v)
543 clause.append(_where_clause(k, expr, q, metadata_t, embeddings_t))
544 return reduce(lambda x, y: x & y, clause)
545
546 @trace_method(
547 "SqliteMetadataSegment._where_doc_criterion", OpenTelemetryGranularity.ALL
548 )
549 def _where_doc_criterion(
550 self,
551 q: QueryBuilder,
552 where: WhereDocument,
553 metadata_t: Table,
554 fulltext_t: Table,
555 embeddings_t: Table,
556 ) -> Criterion:
557 for k, v in where.items():
558 if k == "$and":
559 criteria = [
560 self._where_doc_criterion(
561 q, w, metadata_t, fulltext_t, embeddings_t
562 )
563 for w in cast(Sequence[WhereDocument], v)
564 ]
565 return reduce(lambda x, y: x & y, criteria)
566 elif k == "$or":
567 criteria = [
568 self._where_doc_criterion(
569 q, w, metadata_t, fulltext_t, embeddings_t
570 )
571 for w in cast(Sequence[WhereDocument], v)
572 ]
573 return reduce(lambda x, y: x | y, criteria)
574 elif k in ("$contains", "$not_contains"):
575 v = cast(str, v)
576 search_term = f"%{v}%"
577
578 sq = (
579 self._db.querybuilder()
580 .from_(fulltext_t)
581 .select(fulltext_t.rowid)
582 .where(fulltext_t.string_value.like(ParameterValue(search_term)))
583 )
584 return (
585 embeddings_t.id.isin(sq)
586 if k == "$contains"
587 else embeddings_t.id.notin(sq)
588 )
589 else:
590 raise ValueError(f"Unknown where_doc operator {k}")
591 raise ValueError("Empty where_doc")
592
593 @trace_method("SqliteMetadataSegment.delete", OpenTelemetryGranularity.ALL)
594 @override
595 def delete(self) -> None:
596 t = Table("embeddings")
597 t1 = Table("embedding_metadata")
598 t2 = Table("embedding_fulltext_search")
599 q0 = (
600 self._db.querybuilder()
601 .from_(t1)
602 .delete()
603 .where(
604 t1.id.isin(
605 self._db.querybuilder()
606 .from_(t)
607 .select(t.id)
608 .where(
609 t.segment_id == ParameterValue(self._db.uuid_to_db(self._id))
610 )
611 )
612 )
613 )
614 q = (
615 self._db.querybuilder()
616 .from_(t)
617 .delete()
618 .where(
619 t.id.isin(
620 self._db.querybuilder()
621 .from_(t)
622 .select(t.id)
623 .where(
624 t.segment_id == ParameterValue(self._db.uuid_to_db(self._id))
625 )
626 )
627 )
628 )
629 q_fts = (
630 self._db.querybuilder()
631 .from_(t2)
632 .delete()
633 .where(
634 t2.rowid.isin(
635 self._db.querybuilder()
636 .from_(t)
637 .select(t.id)
638 .where(
639 t.segment_id == ParameterValue(self._db.uuid_to_db(self._id))
640 )
641 )
642 )
643 )
644 with self._db.tx() as cur:
645 cur.execute(*get_sql(q_fts))
646 cur.execute(*get_sql(q0))
647 cur.execute(*get_sql(q))
648
649
650def _where_clause(
651 key: str,
652 expr: Union[
653 LiteralValue,
654 Dict[WhereOperator, LiteralValue],
655 Dict[InclusionExclusionOperator, List[LiteralValue]],
656 ],
657 metadata_q: QueryBuilder,
658 metadata_t: Table,
659 embeddings_t: Table,
660) -> Criterion:
661 """Given a field name, an expression, and a table, construct a Pypika Criterion"""
662
663 # Literal value case
664 if isinstance(expr, (str, int, float, bool)):
665 return _where_clause(
666 key,
667 {cast(WhereOperator, "$eq"): expr},
668 metadata_q,
669 metadata_t,
670 embeddings_t,
671 )
672
673 # Operator dict case
674 operator, value = next(iter(expr.items()))
675 return _value_criterion(key, value, operator, metadata_q, metadata_t, embeddings_t)
676
677
678def _value_criterion(
679 key: str,
680 value: Union[LiteralValue, List[LiteralValue]],
681 op: Union[WhereOperator, InclusionExclusionOperator],
682 metadata_q: QueryBuilder,
683 metadata_t: Table,
684 embeddings_t: Table,
685) -> Criterion:
686 """Creates the filter for a single operator"""
687
688 def is_numeric(obj: object) -> bool:
689 return (not isinstance(obj, bool)) and isinstance(obj, (int, float))
690
691 sub_q = metadata_q.where(metadata_t.key == ParameterValue(key))
692 p_val = ParameterValue(value)
693
694 if is_numeric(value) or (isinstance(value, list) and is_numeric(value[0])):
695 int_col, float_col = metadata_t.int_value, metadata_t.float_value
696 if op in ("$eq", "$ne"):
697 expr = (int_col == p_val) | (float_col == p_val)
698 elif op == "$gt":
699 expr = (int_col > p_val) | (float_col > p_val)
700 elif op == "$gte":
701 expr = (int_col >= p_val) | (float_col >= p_val)
702 elif op == "$lt":
703 expr = (int_col < p_val) | (float_col < p_val)
704 elif op == "$lte":
705 expr = (int_col <= p_val) | (float_col <= p_val)
706 else:
707 expr = int_col.isin(p_val) | float_col.isin(p_val)
708 else:
709 if isinstance(value, bool) or (
710 isinstance(value, list) and isinstance(value[0], bool)
711 ):
712 col = metadata_t.bool_value
713 else:
714 col = metadata_t.string_value
715 if op in ("$eq", "$ne"):
716 expr = col == p_val
717 else:
718 expr = col.isin(p_val)
719
720 if op in ("$ne", "$nin"):
721 return embeddings_t.id.notin(sub_q.where(expr))
722 else:
723 return embeddings_t.id.isin(sub_q.where(expr))
724 