Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes15kdownloads
sqlite.py724 linesDownload Raw Back to metadata
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 
codekingpro/portable-devtools · Team Ai