chore: 添加虚拟环境到仓库
- 添加 backend_service/venv 虚拟环境 - 包含所有Python依赖包 - 注意:虚拟环境约393MB,包含12655个文件
This commit is contained in:
@@ -0,0 +1,507 @@
|
||||
from functools import cached_property
|
||||
import json
|
||||
from chromadb.api.configuration import (
|
||||
ConfigurationParameter,
|
||||
EmbeddingsQueueConfigurationInternal,
|
||||
)
|
||||
from chromadb.db.base import SqlDB, ParameterValue, get_sql
|
||||
from chromadb.errors import BatchSizeExceededError
|
||||
from chromadb.ingest import (
|
||||
Producer,
|
||||
Consumer,
|
||||
ConsumerCallbackFn,
|
||||
decode_vector,
|
||||
encode_vector,
|
||||
)
|
||||
from chromadb.types import (
|
||||
OperationRecord,
|
||||
LogRecord,
|
||||
ScalarEncoding,
|
||||
SeqId,
|
||||
Operation,
|
||||
)
|
||||
from chromadb.config import System
|
||||
from chromadb.telemetry.opentelemetry import (
|
||||
OpenTelemetryClient,
|
||||
OpenTelemetryGranularity,
|
||||
trace_method,
|
||||
)
|
||||
from overrides import override
|
||||
from collections import defaultdict
|
||||
from typing import Sequence, Optional, Dict, Set, Tuple, cast
|
||||
from uuid import UUID
|
||||
from pypika import Table, functions
|
||||
import uuid
|
||||
import logging
|
||||
from chromadb.ingest.impl.utils import create_topic_name
|
||||
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_operation_codes = {
|
||||
Operation.ADD: 0,
|
||||
Operation.UPDATE: 1,
|
||||
Operation.UPSERT: 2,
|
||||
Operation.DELETE: 3,
|
||||
}
|
||||
_operation_codes_inv = {v: k for k, v in _operation_codes.items()}
|
||||
|
||||
# Set in conftest.py to rethrow errors in the "async" path during testing
|
||||
# https://doc.pytest.org/en/latest/example/simple.html#detect-if-running-from-within-a-pytest-run
|
||||
_called_from_test = False
|
||||
|
||||
|
||||
class SqlEmbeddingsQueue(SqlDB, Producer, Consumer):
|
||||
"""A SQL database that stores embeddings, allowing a traditional RDBMS to be used as
|
||||
the primary ingest queue and satisfying the top level Producer/Consumer interfaces.
|
||||
|
||||
Note that this class is only suitable for use cases where the producer and consumer
|
||||
are in the same process.
|
||||
|
||||
This is because notification of new embeddings happens solely in-process: this
|
||||
implementation does not actively listen to the the database for new records added by
|
||||
other processes.
|
||||
"""
|
||||
|
||||
class Subscription:
|
||||
id: UUID
|
||||
topic_name: str
|
||||
start: int
|
||||
end: int
|
||||
callback: ConsumerCallbackFn
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
id: UUID,
|
||||
topic_name: str,
|
||||
start: int,
|
||||
end: int,
|
||||
callback: ConsumerCallbackFn,
|
||||
):
|
||||
self.id = id
|
||||
self.topic_name = topic_name
|
||||
self.start = start
|
||||
self.end = end
|
||||
self.callback = callback
|
||||
|
||||
_subscriptions: Dict[str, Set[Subscription]]
|
||||
_max_batch_size: Optional[int]
|
||||
_tenant: str
|
||||
_topic_namespace: str
|
||||
# How many variables are in the insert statement for a single record
|
||||
VARIABLES_PER_RECORD = 6
|
||||
|
||||
def __init__(self, system: System):
|
||||
self._subscriptions = defaultdict(set)
|
||||
self._max_batch_size = None
|
||||
self._opentelemetry_client = system.require(OpenTelemetryClient)
|
||||
self._tenant = system.settings.require("tenant_id")
|
||||
self._topic_namespace = system.settings.require("topic_namespace")
|
||||
super().__init__(system)
|
||||
|
||||
@trace_method("SqlEmbeddingsQueue.reset_state", OpenTelemetryGranularity.ALL)
|
||||
@override
|
||||
def reset_state(self) -> None:
|
||||
super().reset_state()
|
||||
self._subscriptions = defaultdict(set)
|
||||
|
||||
# Invalidate the cached property
|
||||
try:
|
||||
del self.config
|
||||
except AttributeError:
|
||||
# Cached property hasn't been accessed yet
|
||||
pass
|
||||
|
||||
@trace_method("SqlEmbeddingsQueue.delete_topic", OpenTelemetryGranularity.ALL)
|
||||
@override
|
||||
def delete_log(self, collection_id: UUID) -> None:
|
||||
topic_name = create_topic_name(
|
||||
self._tenant, self._topic_namespace, collection_id
|
||||
)
|
||||
t = Table("embeddings_queue")
|
||||
q = (
|
||||
self.querybuilder()
|
||||
.from_(t)
|
||||
.where(t.topic == ParameterValue(topic_name))
|
||||
.delete()
|
||||
)
|
||||
with self.tx() as cur:
|
||||
sql, params = get_sql(q, self.parameter_format())
|
||||
cur.execute(sql, params)
|
||||
|
||||
@trace_method("SqlEmbeddingsQueue.purge_log", OpenTelemetryGranularity.ALL)
|
||||
@override
|
||||
def purge_log(self, collection_id: UUID) -> None:
|
||||
# (We need to purge on a per topic/collection basis, because the maximum sequence ID is tracked on a per topic/collection basis.)
|
||||
|
||||
segments_t = Table("segments")
|
||||
segment_ids_q = (
|
||||
self.querybuilder()
|
||||
.from_(segments_t)
|
||||
# This coalesce prevents a correctness bug when > 1 segments exist and:
|
||||
# - > 1 has written to the max_seq_id table
|
||||
# - > 1 has not never written to the max_seq_id table
|
||||
# In that case, we should not delete any WAL entries as we can't be sure that the all segments are caught up.
|
||||
.select(functions.Coalesce(Table("max_seq_id").seq_id, -1))
|
||||
.where(
|
||||
segments_t.collection == ParameterValue(self.uuid_to_db(collection_id))
|
||||
)
|
||||
.left_join(Table("max_seq_id"))
|
||||
.on(segments_t.id == Table("max_seq_id").segment_id)
|
||||
)
|
||||
|
||||
topic_name = create_topic_name(
|
||||
self._tenant, self._topic_namespace, collection_id
|
||||
)
|
||||
with self.tx() as cur:
|
||||
sql, params = get_sql(segment_ids_q, self.parameter_format())
|
||||
cur.execute(sql, params)
|
||||
results = cur.fetchall()
|
||||
if results:
|
||||
min_seq_id = min(row[0] for row in results)
|
||||
else:
|
||||
return
|
||||
|
||||
t = Table("embeddings_queue")
|
||||
q = (
|
||||
self.querybuilder()
|
||||
.from_(t)
|
||||
.where(t.seq_id < ParameterValue(min_seq_id))
|
||||
.where(t.topic == ParameterValue(topic_name))
|
||||
.delete()
|
||||
)
|
||||
|
||||
sql, params = get_sql(q, self.parameter_format())
|
||||
cur.execute(sql, params)
|
||||
|
||||
@trace_method("SqlEmbeddingsQueue.submit_embedding", OpenTelemetryGranularity.ALL)
|
||||
@override
|
||||
def submit_embedding(
|
||||
self, collection_id: UUID, embedding: OperationRecord
|
||||
) -> SeqId:
|
||||
if not self._running:
|
||||
raise RuntimeError("Component not running")
|
||||
|
||||
return self.submit_embeddings(collection_id, [embedding])[0]
|
||||
|
||||
@trace_method("SqlEmbeddingsQueue.submit_embeddings", OpenTelemetryGranularity.ALL)
|
||||
@override
|
||||
def submit_embeddings(
|
||||
self, collection_id: UUID, embeddings: Sequence[OperationRecord]
|
||||
) -> Sequence[SeqId]:
|
||||
if not self._running:
|
||||
raise RuntimeError("Component not running")
|
||||
|
||||
if len(embeddings) == 0:
|
||||
return []
|
||||
|
||||
if len(embeddings) > self.max_batch_size:
|
||||
raise BatchSizeExceededError(
|
||||
f"""
|
||||
Cannot submit more than {self.max_batch_size:,} embeddings at once.
|
||||
Please submit your embeddings in batches of size
|
||||
{self.max_batch_size:,} or less.
|
||||
"""
|
||||
)
|
||||
|
||||
# This creates the persisted configuration if it doesn't exist.
|
||||
# It should be run as soon as possible (before any WAL mutations) since the default configuration depends on the WAL size.
|
||||
# (We can't run this in __init__()/start() because the migrations have not been run at that point and the table may not be available.)
|
||||
_ = self.config
|
||||
|
||||
topic_name = create_topic_name(
|
||||
self._tenant, self._topic_namespace, collection_id
|
||||
)
|
||||
|
||||
t = Table("embeddings_queue")
|
||||
insert = (
|
||||
self.querybuilder()
|
||||
.into(t)
|
||||
.columns(t.operation, t.topic, t.id, t.vector, t.encoding, t.metadata)
|
||||
)
|
||||
id_to_idx: Dict[str, int] = {}
|
||||
for embedding in embeddings:
|
||||
(
|
||||
embedding_bytes,
|
||||
encoding,
|
||||
metadata,
|
||||
) = self._prepare_vector_encoding_metadata(embedding)
|
||||
insert = insert.insert(
|
||||
ParameterValue(_operation_codes[embedding["operation"]]),
|
||||
ParameterValue(topic_name),
|
||||
ParameterValue(embedding["id"]),
|
||||
ParameterValue(embedding_bytes),
|
||||
ParameterValue(encoding),
|
||||
ParameterValue(metadata),
|
||||
)
|
||||
id_to_idx[embedding["id"]] = len(id_to_idx)
|
||||
with self.tx() as cur:
|
||||
sql, params = get_sql(insert, self.parameter_format())
|
||||
# The returning clause does not guarantee order, so we need to do reorder
|
||||
# the results. https://www.sqlite.org/lang_returning.html
|
||||
sql = f"{sql} RETURNING seq_id, id" # Pypika doesn't support RETURNING
|
||||
results = cur.execute(sql, params).fetchall()
|
||||
# Reorder the results
|
||||
seq_ids = [cast(SeqId, None)] * len(
|
||||
results
|
||||
) # Lie to mypy: https://stackoverflow.com/questions/76694215/python-type-casting-when-preallocating-list
|
||||
embedding_records = []
|
||||
for seq_id, id in results:
|
||||
seq_ids[id_to_idx[id]] = seq_id
|
||||
submit_embedding_record = embeddings[id_to_idx[id]]
|
||||
# We allow notifying consumers out of order relative to one call to
|
||||
# submit_embeddings so we do not reorder the records before submitting them
|
||||
embedding_record = LogRecord(
|
||||
log_offset=seq_id,
|
||||
record=OperationRecord(
|
||||
id=id,
|
||||
embedding=submit_embedding_record["embedding"],
|
||||
encoding=submit_embedding_record["encoding"],
|
||||
metadata=submit_embedding_record["metadata"],
|
||||
operation=submit_embedding_record["operation"],
|
||||
),
|
||||
)
|
||||
embedding_records.append(embedding_record)
|
||||
self._notify_all(topic_name, embedding_records)
|
||||
|
||||
if self.config.get_parameter("automatically_purge").value:
|
||||
self.purge_log(collection_id)
|
||||
|
||||
return seq_ids
|
||||
|
||||
@trace_method("SqlEmbeddingsQueue.subscribe", OpenTelemetryGranularity.ALL)
|
||||
@override
|
||||
def subscribe(
|
||||
self,
|
||||
collection_id: UUID,
|
||||
consume_fn: ConsumerCallbackFn,
|
||||
start: Optional[SeqId] = None,
|
||||
end: Optional[SeqId] = None,
|
||||
id: Optional[UUID] = None,
|
||||
) -> UUID:
|
||||
if not self._running:
|
||||
raise RuntimeError("Component not running")
|
||||
|
||||
topic_name = create_topic_name(
|
||||
self._tenant, self._topic_namespace, collection_id
|
||||
)
|
||||
|
||||
subscription_id = id or uuid.uuid4()
|
||||
start, end = self._validate_range(start, end)
|
||||
|
||||
subscription = self.Subscription(
|
||||
subscription_id, topic_name, start, end, consume_fn
|
||||
)
|
||||
|
||||
# Backfill first, so if it errors we do not add the subscription
|
||||
self._backfill(subscription)
|
||||
self._subscriptions[topic_name].add(subscription)
|
||||
|
||||
return subscription_id
|
||||
|
||||
@trace_method("SqlEmbeddingsQueue.unsubscribe", OpenTelemetryGranularity.ALL)
|
||||
@override
|
||||
def unsubscribe(self, subscription_id: UUID) -> None:
|
||||
for topic_name, subscriptions in self._subscriptions.items():
|
||||
for subscription in subscriptions:
|
||||
if subscription.id == subscription_id:
|
||||
subscriptions.remove(subscription)
|
||||
if len(subscriptions) == 0:
|
||||
del self._subscriptions[topic_name]
|
||||
return
|
||||
|
||||
@override
|
||||
def min_seqid(self) -> SeqId:
|
||||
return -1
|
||||
|
||||
@override
|
||||
def max_seqid(self) -> SeqId:
|
||||
return 2**63 - 1
|
||||
|
||||
@property
|
||||
@trace_method("SqlEmbeddingsQueue.max_batch_size", OpenTelemetryGranularity.ALL)
|
||||
@override
|
||||
def max_batch_size(self) -> int:
|
||||
if self._max_batch_size is None:
|
||||
with self.tx() as cur:
|
||||
cur.execute("PRAGMA compile_options;")
|
||||
compile_options = cur.fetchall()
|
||||
|
||||
for option in compile_options:
|
||||
if "MAX_VARIABLE_NUMBER" in option[0]:
|
||||
# The pragma returns a string like 'MAX_VARIABLE_NUMBER=999'
|
||||
self._max_batch_size = int(option[0].split("=")[1]) // (
|
||||
self.VARIABLES_PER_RECORD
|
||||
)
|
||||
|
||||
if self._max_batch_size is None:
|
||||
# This value is the default for sqlite3 versions < 3.32.0
|
||||
# It is the safest value to use if we can't find the pragma for some
|
||||
# reason
|
||||
self._max_batch_size = 999 // self.VARIABLES_PER_RECORD
|
||||
return self._max_batch_size
|
||||
|
||||
@trace_method(
|
||||
"SqlEmbeddingsQueue._prepare_vector_encoding_metadata",
|
||||
OpenTelemetryGranularity.ALL,
|
||||
)
|
||||
def _prepare_vector_encoding_metadata(
|
||||
self, embedding: OperationRecord
|
||||
) -> Tuple[Optional[bytes], Optional[str], Optional[str]]:
|
||||
if embedding["embedding"] is not None:
|
||||
encoding_type = cast(ScalarEncoding, embedding["encoding"])
|
||||
encoding = encoding_type.value
|
||||
embedding_bytes = encode_vector(embedding["embedding"], encoding_type)
|
||||
else:
|
||||
embedding_bytes = None
|
||||
encoding = None
|
||||
metadata = json.dumps(embedding["metadata"]) if embedding["metadata"] else None
|
||||
return embedding_bytes, encoding, metadata
|
||||
|
||||
@trace_method("SqlEmbeddingsQueue._backfill", OpenTelemetryGranularity.ALL)
|
||||
def _backfill(self, subscription: Subscription) -> None:
|
||||
"""Backfill the given subscription with any currently matching records in the
|
||||
DB"""
|
||||
t = Table("embeddings_queue")
|
||||
q = (
|
||||
self.querybuilder()
|
||||
.from_(t)
|
||||
.where(t.topic == ParameterValue(subscription.topic_name))
|
||||
.where(t.seq_id > ParameterValue(subscription.start))
|
||||
.where(t.seq_id <= ParameterValue(subscription.end))
|
||||
.select(t.seq_id, t.operation, t.id, t.vector, t.encoding, t.metadata)
|
||||
.orderby(t.seq_id)
|
||||
)
|
||||
with self.tx() as cur:
|
||||
sql, params = get_sql(q, self.parameter_format())
|
||||
cur.execute(sql, params)
|
||||
rows = cur.fetchall()
|
||||
for row in rows:
|
||||
if row[3]:
|
||||
encoding = ScalarEncoding(row[4])
|
||||
vector = decode_vector(row[3], encoding)
|
||||
else:
|
||||
encoding = None
|
||||
vector = None
|
||||
self._notify_one(
|
||||
subscription,
|
||||
[
|
||||
LogRecord(
|
||||
log_offset=row[0],
|
||||
record=OperationRecord(
|
||||
operation=_operation_codes_inv[row[1]],
|
||||
id=row[2],
|
||||
embedding=vector,
|
||||
encoding=encoding,
|
||||
metadata=json.loads(row[5]) if row[5] else None,
|
||||
),
|
||||
)
|
||||
],
|
||||
)
|
||||
|
||||
@trace_method("SqlEmbeddingsQueue._validate_range", OpenTelemetryGranularity.ALL)
|
||||
def _validate_range(
|
||||
self, start: Optional[SeqId], end: Optional[SeqId]
|
||||
) -> Tuple[int, int]:
|
||||
"""Validate and normalize the start and end SeqIDs for a subscription using this
|
||||
impl."""
|
||||
start = start or self._next_seq_id()
|
||||
end = end or self.max_seqid()
|
||||
if not isinstance(start, int) or not isinstance(end, int):
|
||||
raise TypeError("SeqIDs must be integers for sql-based EmbeddingsDB")
|
||||
if start >= end:
|
||||
raise ValueError(f"Invalid SeqID range: {start} to {end}")
|
||||
return start, end
|
||||
|
||||
@trace_method("SqlEmbeddingsQueue._next_seq_id", OpenTelemetryGranularity.ALL)
|
||||
def _next_seq_id(self) -> int:
|
||||
"""Get the next SeqID for this database."""
|
||||
t = Table("embeddings_queue")
|
||||
q = self.querybuilder().from_(t).select(functions.Max(t.seq_id))
|
||||
with self.tx() as cur:
|
||||
cur.execute(q.get_sql())
|
||||
return int(cur.fetchone()[0]) + 1
|
||||
|
||||
@trace_method("SqlEmbeddingsQueue._notify_all", OpenTelemetryGranularity.ALL)
|
||||
def _notify_all(self, topic: str, embeddings: Sequence[LogRecord]) -> None:
|
||||
"""Send a notification to each subscriber of the given topic."""
|
||||
if self._running:
|
||||
for sub in self._subscriptions[topic]:
|
||||
self._notify_one(sub, embeddings)
|
||||
|
||||
@trace_method("SqlEmbeddingsQueue._notify_one", OpenTelemetryGranularity.ALL)
|
||||
def _notify_one(self, sub: Subscription, embeddings: Sequence[LogRecord]) -> None:
|
||||
"""Send a notification to a single subscriber."""
|
||||
# Filter out any embeddings that are not in the subscription range
|
||||
should_unsubscribe = False
|
||||
filtered_embeddings = []
|
||||
for embedding in embeddings:
|
||||
if embedding["log_offset"] <= sub.start:
|
||||
continue
|
||||
if embedding["log_offset"] > sub.end:
|
||||
should_unsubscribe = True
|
||||
break
|
||||
filtered_embeddings.append(embedding)
|
||||
|
||||
# Log errors instead of throwing them to preserve async semantics
|
||||
# for consistency between local and distributed configurations
|
||||
try:
|
||||
if len(filtered_embeddings) > 0:
|
||||
sub.callback(filtered_embeddings)
|
||||
if should_unsubscribe:
|
||||
self.unsubscribe(sub.id)
|
||||
except BaseException as e:
|
||||
logger.error(
|
||||
f"Exception occurred invoking consumer for subscription {sub.id.hex}"
|
||||
+ f"to topic {sub.topic_name} %s",
|
||||
str(e),
|
||||
)
|
||||
if _called_from_test:
|
||||
raise e
|
||||
|
||||
@cached_property
|
||||
def config(self) -> EmbeddingsQueueConfigurationInternal:
|
||||
t = Table("embeddings_queue_config")
|
||||
q = self.querybuilder().from_(t).select(t.config_json_str).limit(1)
|
||||
|
||||
with self.tx() as cur:
|
||||
cur.execute(q.get_sql())
|
||||
result = cur.fetchone()
|
||||
|
||||
if result is None:
|
||||
is_fresh_system = self._get_wal_size() == 0
|
||||
config = EmbeddingsQueueConfigurationInternal(
|
||||
[ConfigurationParameter("automatically_purge", is_fresh_system)]
|
||||
)
|
||||
self.set_config(config)
|
||||
return config
|
||||
|
||||
return EmbeddingsQueueConfigurationInternal.from_json_str(result[0])
|
||||
|
||||
def set_config(self, config: EmbeddingsQueueConfigurationInternal) -> None:
|
||||
with self.tx() as cur:
|
||||
cur.execute(
|
||||
"""
|
||||
INSERT OR REPLACE INTO embeddings_queue_config (id, config_json_str)
|
||||
VALUES (?, ?)
|
||||
""",
|
||||
(
|
||||
1,
|
||||
config.to_json_str(),
|
||||
),
|
||||
)
|
||||
|
||||
# Invalidate the cached property
|
||||
try:
|
||||
del self.config
|
||||
except AttributeError:
|
||||
# Cached property hasn't been accessed yet
|
||||
pass
|
||||
|
||||
def _get_wal_size(self) -> int:
|
||||
t = Table("embeddings_queue")
|
||||
q = self.querybuilder().from_(t).select(functions.Count("*"))
|
||||
|
||||
with self.tx() as cur:
|
||||
cur.execute(q.get_sql())
|
||||
return int(cur.fetchone()[0])
|
||||
@@ -0,0 +1,986 @@
|
||||
import logging
|
||||
import sys
|
||||
from typing import Optional, Sequence, Any, Tuple, cast, Dict, Union, Set
|
||||
from uuid import UUID
|
||||
from overrides import override
|
||||
from pypika import Table, Column
|
||||
from itertools import groupby
|
||||
|
||||
from chromadb.api.types import Schema
|
||||
from chromadb.config import DEFAULT_DATABASE, DEFAULT_TENANT, System
|
||||
from chromadb.db.base import Cursor, SqlDB, ParameterValue, get_sql
|
||||
from chromadb.db.system import SysDB
|
||||
from chromadb.errors import (
|
||||
NotFoundError,
|
||||
UniqueConstraintError,
|
||||
)
|
||||
from chromadb.telemetry.opentelemetry import (
|
||||
add_attributes_to_current_span,
|
||||
OpenTelemetryClient,
|
||||
OpenTelemetryGranularity,
|
||||
trace_method,
|
||||
)
|
||||
from chromadb.ingest import Producer
|
||||
from chromadb.types import (
|
||||
CollectionAndSegments,
|
||||
Database,
|
||||
OptionalArgument,
|
||||
Segment,
|
||||
Metadata,
|
||||
Collection,
|
||||
SegmentScope,
|
||||
Tenant,
|
||||
Unspecified,
|
||||
UpdateMetadata,
|
||||
)
|
||||
from chromadb.api.collection_configuration import (
|
||||
CreateCollectionConfiguration,
|
||||
UpdateCollectionConfiguration,
|
||||
create_collection_configuration_to_json_str,
|
||||
load_collection_configuration_from_json_str,
|
||||
CollectionConfiguration,
|
||||
create_collection_configuration_to_json,
|
||||
collection_configuration_to_json,
|
||||
collection_configuration_to_json_str,
|
||||
overwrite_collection_configuration,
|
||||
update_collection_configuration_from_legacy_update_metadata,
|
||||
CollectionMetadata,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class SqlSysDB(SqlDB, SysDB):
|
||||
# Used only to delete log streams on collection deletion.
|
||||
# TODO: refactor to remove this dependency into a separate interface
|
||||
_producer: Producer
|
||||
|
||||
def __init__(self, system: System):
|
||||
super().__init__(system)
|
||||
self._opentelemetry_client = system.require(OpenTelemetryClient)
|
||||
|
||||
@trace_method("SqlSysDB.create_segment", OpenTelemetryGranularity.ALL)
|
||||
@override
|
||||
def start(self) -> None:
|
||||
super().start()
|
||||
self._producer = self._system.instance(Producer)
|
||||
|
||||
@override
|
||||
def create_database(
|
||||
self, id: UUID, name: str, tenant: str = DEFAULT_TENANT
|
||||
) -> None:
|
||||
with self.tx() as cur:
|
||||
# Get the tenant id for the tenant name and then insert the database with the id, name and tenant id
|
||||
databases = Table("databases")
|
||||
tenants = Table("tenants")
|
||||
insert_database = (
|
||||
self.querybuilder()
|
||||
.into(databases)
|
||||
.columns(databases.id, databases.name, databases.tenant_id)
|
||||
.insert(
|
||||
ParameterValue(self.uuid_to_db(id)),
|
||||
ParameterValue(name),
|
||||
self.querybuilder()
|
||||
.select(tenants.id)
|
||||
.from_(tenants)
|
||||
.where(tenants.id == ParameterValue(tenant)),
|
||||
)
|
||||
)
|
||||
sql, params = get_sql(insert_database, self.parameter_format())
|
||||
try:
|
||||
cur.execute(sql, params)
|
||||
except self.unique_constraint_error() as e:
|
||||
raise UniqueConstraintError(
|
||||
f"Database {name} already exists for tenant {tenant}"
|
||||
) from e
|
||||
|
||||
@override
|
||||
def get_database(self, name: str, tenant: str = DEFAULT_TENANT) -> Database:
|
||||
with self.tx() as cur:
|
||||
databases = Table("databases")
|
||||
q = (
|
||||
self.querybuilder()
|
||||
.from_(databases)
|
||||
.select(databases.id, databases.name)
|
||||
.where(databases.name == ParameterValue(name))
|
||||
.where(databases.tenant_id == ParameterValue(tenant))
|
||||
)
|
||||
sql, params = get_sql(q, self.parameter_format())
|
||||
row = cur.execute(sql, params).fetchone()
|
||||
if not row:
|
||||
raise NotFoundError(
|
||||
f"Database {name} not found for tenant {tenant}. Are you sure it exists?"
|
||||
)
|
||||
if row[0] is None:
|
||||
raise NotFoundError(
|
||||
f"Database {name} not found for tenant {tenant}. Are you sure it exists?"
|
||||
)
|
||||
id: UUID = cast(UUID, self.uuid_from_db(row[0]))
|
||||
return Database(
|
||||
id=id,
|
||||
name=row[1],
|
||||
tenant=tenant,
|
||||
)
|
||||
|
||||
@override
|
||||
def delete_database(self, name: str, tenant: str = DEFAULT_TENANT) -> None:
|
||||
with self.tx() as cur:
|
||||
databases = Table("databases")
|
||||
q = (
|
||||
self.querybuilder()
|
||||
.from_(databases)
|
||||
.where(databases.name == ParameterValue(name))
|
||||
.where(databases.tenant_id == ParameterValue(tenant))
|
||||
.delete()
|
||||
)
|
||||
sql, params = get_sql(q, self.parameter_format())
|
||||
sql = sql + " RETURNING id"
|
||||
result = cur.execute(sql, params).fetchone()
|
||||
if not result:
|
||||
raise NotFoundError(f"Database {name} not found for tenant {tenant}")
|
||||
|
||||
# As of 01/09/2025, cascading deletes don't work because foreign keys are not enabled.
|
||||
# See https://github.com/chroma-core/chroma/issues/3456.
|
||||
collections = Table("collections")
|
||||
q = (
|
||||
self.querybuilder()
|
||||
.from_(collections)
|
||||
.where(collections.database_id == ParameterValue(result[0]))
|
||||
.delete()
|
||||
)
|
||||
sql, params = get_sql(q, self.parameter_format())
|
||||
cur.execute(sql, params)
|
||||
|
||||
@override
|
||||
def list_databases(
|
||||
self,
|
||||
limit: Optional[int] = None,
|
||||
offset: Optional[int] = None,
|
||||
tenant: str = DEFAULT_TENANT,
|
||||
) -> Sequence[Database]:
|
||||
with self.tx() as cur:
|
||||
databases = Table("databases")
|
||||
q = (
|
||||
self.querybuilder()
|
||||
.from_(databases)
|
||||
.select(databases.id, databases.name)
|
||||
.where(databases.tenant_id == ParameterValue(tenant))
|
||||
.offset(offset)
|
||||
.limit(
|
||||
sys.maxsize if limit is None else limit
|
||||
) # SQLite requires that a limit is provided to use offset
|
||||
.orderby(databases.created_at)
|
||||
)
|
||||
sql, params = get_sql(q, self.parameter_format())
|
||||
rows = cur.execute(sql, params).fetchall()
|
||||
return [
|
||||
Database(
|
||||
id=cast(UUID, self.uuid_from_db(row[0])),
|
||||
name=row[1],
|
||||
tenant=tenant,
|
||||
)
|
||||
for row in rows
|
||||
]
|
||||
|
||||
@override
|
||||
def create_tenant(self, name: str) -> None:
|
||||
with self.tx() as cur:
|
||||
tenants = Table("tenants")
|
||||
insert_tenant = (
|
||||
self.querybuilder()
|
||||
.into(tenants)
|
||||
.columns(tenants.id)
|
||||
.insert(ParameterValue(name))
|
||||
)
|
||||
sql, params = get_sql(insert_tenant, self.parameter_format())
|
||||
try:
|
||||
cur.execute(sql, params)
|
||||
except self.unique_constraint_error() as e:
|
||||
raise UniqueConstraintError(f"Tenant {name} already exists") from e
|
||||
|
||||
@override
|
||||
def get_tenant(self, name: str) -> Tenant:
|
||||
with self.tx() as cur:
|
||||
tenants = Table("tenants")
|
||||
q = (
|
||||
self.querybuilder()
|
||||
.from_(tenants)
|
||||
.select(tenants.id)
|
||||
.where(tenants.id == ParameterValue(name))
|
||||
)
|
||||
sql, params = get_sql(q, self.parameter_format())
|
||||
row = cur.execute(sql, params).fetchone()
|
||||
if not row:
|
||||
raise NotFoundError(f"Tenant {name} not found")
|
||||
return Tenant(name=name)
|
||||
|
||||
# Create a segment using the passed cursor, so that the other changes
|
||||
# can be in the same transaction.
|
||||
def create_segment_with_tx(self, cur: Cursor, segment: Segment) -> None:
|
||||
add_attributes_to_current_span(
|
||||
{
|
||||
"segment_id": str(segment["id"]),
|
||||
"segment_type": segment["type"],
|
||||
"segment_scope": segment["scope"].value,
|
||||
"collection": str(segment["collection"]),
|
||||
}
|
||||
)
|
||||
|
||||
segments = Table("segments")
|
||||
insert_segment = (
|
||||
self.querybuilder()
|
||||
.into(segments)
|
||||
.columns(
|
||||
segments.id,
|
||||
segments.type,
|
||||
segments.scope,
|
||||
segments.collection,
|
||||
)
|
||||
.insert(
|
||||
ParameterValue(self.uuid_to_db(segment["id"])),
|
||||
ParameterValue(segment["type"]),
|
||||
ParameterValue(segment["scope"].value),
|
||||
ParameterValue(self.uuid_to_db(segment["collection"])),
|
||||
)
|
||||
)
|
||||
sql, params = get_sql(insert_segment, self.parameter_format())
|
||||
try:
|
||||
cur.execute(sql, params)
|
||||
except self.unique_constraint_error() as e:
|
||||
raise UniqueConstraintError(
|
||||
f"Segment {segment['id']} already exists"
|
||||
) from e
|
||||
|
||||
# Insert segment metadata if it exists
|
||||
metadata_t = Table("segment_metadata")
|
||||
if segment["metadata"]:
|
||||
try:
|
||||
self._insert_metadata(
|
||||
cur,
|
||||
metadata_t,
|
||||
metadata_t.segment_id,
|
||||
segment["id"],
|
||||
segment["metadata"],
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(f"Error inserting segment metadata: {e}")
|
||||
raise
|
||||
|
||||
# TODO(rohit): Investigate and remove this method completely.
|
||||
@trace_method("SqlSysDB.create_segment", OpenTelemetryGranularity.ALL)
|
||||
@override
|
||||
def create_segment(self, segment: Segment) -> None:
|
||||
with self.tx() as cur:
|
||||
self.create_segment_with_tx(cur, segment)
|
||||
|
||||
@trace_method("SqlSysDB.create_collection", OpenTelemetryGranularity.ALL)
|
||||
@override
|
||||
def create_collection(
|
||||
self,
|
||||
id: UUID,
|
||||
name: str,
|
||||
schema: Optional[Schema],
|
||||
configuration: CreateCollectionConfiguration,
|
||||
segments: Sequence[Segment],
|
||||
metadata: Optional[Metadata] = None,
|
||||
dimension: Optional[int] = None,
|
||||
get_or_create: bool = False,
|
||||
tenant: str = DEFAULT_TENANT,
|
||||
database: str = DEFAULT_DATABASE,
|
||||
) -> Tuple[Collection, bool]:
|
||||
if id is None and not get_or_create:
|
||||
raise ValueError("id must be specified if get_or_create is False")
|
||||
|
||||
add_attributes_to_current_span(
|
||||
{
|
||||
"collection_id": str(id),
|
||||
"collection_name": name,
|
||||
}
|
||||
)
|
||||
|
||||
existing = self.get_collections(name=name, tenant=tenant, database=database)
|
||||
if existing:
|
||||
if get_or_create:
|
||||
collection = existing[0]
|
||||
return (
|
||||
self.get_collections(
|
||||
id=collection.id, tenant=tenant, database=database
|
||||
)[0],
|
||||
False,
|
||||
)
|
||||
else:
|
||||
raise UniqueConstraintError(f"Collection {name} already exists")
|
||||
|
||||
collection = Collection(
|
||||
id=id,
|
||||
name=name,
|
||||
configuration_json=create_collection_configuration_to_json(
|
||||
configuration, cast(CollectionMetadata, metadata)
|
||||
),
|
||||
serialized_schema=None,
|
||||
metadata=metadata,
|
||||
dimension=dimension,
|
||||
tenant=tenant,
|
||||
database=database,
|
||||
version=0,
|
||||
)
|
||||
|
||||
with self.tx() as cur:
|
||||
collections = Table("collections")
|
||||
databases = Table("databases")
|
||||
|
||||
insert_collection = (
|
||||
self.querybuilder()
|
||||
.into(collections)
|
||||
.columns(
|
||||
collections.id,
|
||||
collections.name,
|
||||
collections.config_json_str,
|
||||
collections.dimension,
|
||||
collections.database_id,
|
||||
)
|
||||
.insert(
|
||||
ParameterValue(self.uuid_to_db(collection["id"])),
|
||||
ParameterValue(collection["name"]),
|
||||
ParameterValue(
|
||||
create_collection_configuration_to_json_str(
|
||||
configuration, cast(CollectionMetadata, metadata)
|
||||
)
|
||||
),
|
||||
ParameterValue(collection["dimension"]),
|
||||
# Get the database id for the database with the given name and tenant
|
||||
self.querybuilder()
|
||||
.select(databases.id)
|
||||
.from_(databases)
|
||||
.where(databases.name == ParameterValue(database))
|
||||
.where(databases.tenant_id == ParameterValue(tenant)),
|
||||
)
|
||||
)
|
||||
sql, params = get_sql(insert_collection, self.parameter_format())
|
||||
try:
|
||||
cur.execute(sql, params)
|
||||
except self.unique_constraint_error() as e:
|
||||
raise UniqueConstraintError(
|
||||
f"Collection {collection['id']} already exists"
|
||||
) from e
|
||||
metadata_t = Table("collection_metadata")
|
||||
if collection["metadata"]:
|
||||
self._insert_metadata(
|
||||
cur,
|
||||
metadata_t,
|
||||
metadata_t.collection_id,
|
||||
collection.id,
|
||||
collection["metadata"],
|
||||
)
|
||||
|
||||
for segment in segments:
|
||||
self.create_segment_with_tx(cur, segment)
|
||||
|
||||
return collection, True
|
||||
|
||||
@trace_method("SqlSysDB.get_segments", OpenTelemetryGranularity.ALL)
|
||||
@override
|
||||
def get_segments(
|
||||
self,
|
||||
collection: UUID,
|
||||
id: Optional[UUID] = None,
|
||||
type: Optional[str] = None,
|
||||
scope: Optional[SegmentScope] = None,
|
||||
) -> Sequence[Segment]:
|
||||
add_attributes_to_current_span(
|
||||
{
|
||||
"segment_id": str(id),
|
||||
"segment_type": type if type else "",
|
||||
"segment_scope": scope.value if scope else "",
|
||||
"collection": str(collection),
|
||||
}
|
||||
)
|
||||
segments_t = Table("segments")
|
||||
metadata_t = Table("segment_metadata")
|
||||
q = (
|
||||
self.querybuilder()
|
||||
.from_(segments_t)
|
||||
.select(
|
||||
segments_t.id,
|
||||
segments_t.type,
|
||||
segments_t.scope,
|
||||
segments_t.collection,
|
||||
metadata_t.key,
|
||||
metadata_t.str_value,
|
||||
metadata_t.int_value,
|
||||
metadata_t.float_value,
|
||||
metadata_t.bool_value,
|
||||
)
|
||||
.left_join(metadata_t)
|
||||
.on(segments_t.id == metadata_t.segment_id)
|
||||
.orderby(segments_t.id)
|
||||
)
|
||||
if id:
|
||||
q = q.where(segments_t.id == ParameterValue(self.uuid_to_db(id)))
|
||||
if type:
|
||||
q = q.where(segments_t.type == ParameterValue(type))
|
||||
if scope:
|
||||
q = q.where(segments_t.scope == ParameterValue(scope.value))
|
||||
if collection:
|
||||
q = q.where(
|
||||
segments_t.collection == ParameterValue(self.uuid_to_db(collection))
|
||||
)
|
||||
|
||||
with self.tx() as cur:
|
||||
sql, params = get_sql(q, self.parameter_format())
|
||||
rows = cur.execute(sql, params).fetchall()
|
||||
by_segment = groupby(rows, lambda r: cast(object, r[0]))
|
||||
segments = []
|
||||
for segment_id, segment_rows in by_segment:
|
||||
id = self.uuid_from_db(str(segment_id))
|
||||
rows = list(segment_rows)
|
||||
type = str(rows[0][1])
|
||||
scope = SegmentScope(str(rows[0][2]))
|
||||
collection = self.uuid_from_db(rows[0][3]) # type: ignore[assignment]
|
||||
metadata = self._metadata_from_rows(rows)
|
||||
segments.append(
|
||||
Segment(
|
||||
id=cast(UUID, id),
|
||||
type=type,
|
||||
scope=scope,
|
||||
collection=collection,
|
||||
metadata=metadata,
|
||||
file_paths={},
|
||||
)
|
||||
)
|
||||
|
||||
return segments
|
||||
|
||||
@trace_method("SqlSysDB.get_collections", OpenTelemetryGranularity.ALL)
|
||||
@override
|
||||
def get_collections(
|
||||
self,
|
||||
id: Optional[UUID] = None,
|
||||
name: Optional[str] = None,
|
||||
tenant: str = DEFAULT_TENANT,
|
||||
database: str = DEFAULT_DATABASE,
|
||||
limit: Optional[int] = None,
|
||||
offset: Optional[int] = None,
|
||||
) -> Sequence[Collection]:
|
||||
"""Get collections by name, embedding function and/or metadata"""
|
||||
|
||||
if name is not None and (tenant is None or database is None):
|
||||
raise ValueError(
|
||||
"If name is specified, tenant and database must also be specified in order to uniquely identify the collection"
|
||||
)
|
||||
|
||||
add_attributes_to_current_span(
|
||||
{
|
||||
"collection_id": str(id),
|
||||
"collection_name": name if name else "",
|
||||
}
|
||||
)
|
||||
|
||||
collections_t = Table("collections")
|
||||
metadata_t = Table("collection_metadata")
|
||||
databases_t = Table("databases")
|
||||
q = (
|
||||
self.querybuilder()
|
||||
.from_(collections_t)
|
||||
.select(
|
||||
collections_t.id,
|
||||
collections_t.name,
|
||||
collections_t.config_json_str,
|
||||
collections_t.dimension,
|
||||
databases_t.name,
|
||||
databases_t.tenant_id,
|
||||
metadata_t.key,
|
||||
metadata_t.str_value,
|
||||
metadata_t.int_value,
|
||||
metadata_t.float_value,
|
||||
metadata_t.bool_value,
|
||||
)
|
||||
.left_join(metadata_t)
|
||||
.on(collections_t.id == metadata_t.collection_id)
|
||||
.left_join(databases_t)
|
||||
.on(collections_t.database_id == databases_t.id)
|
||||
.orderby(collections_t.id)
|
||||
)
|
||||
if id:
|
||||
q = q.where(collections_t.id == ParameterValue(self.uuid_to_db(id)))
|
||||
if name:
|
||||
q = q.where(collections_t.name == ParameterValue(name))
|
||||
|
||||
# Only if we have a name, tenant and database do we need to filter databases
|
||||
# Given an id, we can uniquely identify the collection so we don't need to filter databases
|
||||
if id is None and tenant and database:
|
||||
databases_t = Table("databases")
|
||||
q = q.where(
|
||||
collections_t.database_id
|
||||
== self.querybuilder()
|
||||
.select(databases_t.id)
|
||||
.from_(databases_t)
|
||||
.where(databases_t.name == ParameterValue(database))
|
||||
.where(databases_t.tenant_id == ParameterValue(tenant))
|
||||
)
|
||||
# cant set limit and offset here because this is metadata and we havent reduced yet
|
||||
|
||||
with self.tx() as cur:
|
||||
sql, params = get_sql(q, self.parameter_format())
|
||||
rows = cur.execute(sql, params).fetchall()
|
||||
by_collection = groupby(rows, lambda r: cast(object, r[0]))
|
||||
collections = []
|
||||
for collection_id, collection_rows in by_collection:
|
||||
id = self.uuid_from_db(str(collection_id))
|
||||
rows = list(collection_rows)
|
||||
name = str(rows[0][1])
|
||||
metadata = self._metadata_from_rows(rows)
|
||||
dimension = int(rows[0][3]) if rows[0][3] else None
|
||||
if rows[0][2] is not None:
|
||||
configuration = load_collection_configuration_from_json_str(
|
||||
rows[0][2]
|
||||
)
|
||||
else:
|
||||
# 07/2024: This is a legacy case where we don't have a collection
|
||||
# configuration stored in the database. This non-destructively migrates
|
||||
# the collection to have a configuration, and takes into account any
|
||||
# HNSW params that might be in the existing metadata.
|
||||
configuration = self._insert_config_from_legacy_params(
|
||||
collection_id, metadata
|
||||
)
|
||||
|
||||
collections.append(
|
||||
Collection(
|
||||
id=cast(UUID, id),
|
||||
name=name,
|
||||
configuration_json=collection_configuration_to_json(
|
||||
configuration
|
||||
),
|
||||
serialized_schema=None,
|
||||
metadata=metadata,
|
||||
dimension=dimension,
|
||||
tenant=str(rows[0][5]),
|
||||
database=str(rows[0][4]),
|
||||
version=0,
|
||||
)
|
||||
)
|
||||
|
||||
# apply limit and offset
|
||||
if limit is not None:
|
||||
if offset is None:
|
||||
offset = 0
|
||||
collections = collections[offset : offset + limit]
|
||||
else:
|
||||
collections = collections[offset:]
|
||||
|
||||
return collections
|
||||
|
||||
@override
|
||||
def get_collection_with_segments(
|
||||
self, collection_id: UUID
|
||||
) -> CollectionAndSegments:
|
||||
collections = self.get_collections(id=collection_id)
|
||||
if len(collections) == 0:
|
||||
raise NotFoundError(f"Collection {collection_id} does not exist.")
|
||||
return CollectionAndSegments(
|
||||
collection=collections[0],
|
||||
segments=self.get_segments(collection=collection_id),
|
||||
)
|
||||
|
||||
@trace_method("SqlSysDB.delete_segment", OpenTelemetryGranularity.ALL)
|
||||
@override
|
||||
def delete_segment(self, collection: UUID, id: UUID) -> None:
|
||||
"""Delete a segment from the SysDB"""
|
||||
add_attributes_to_current_span(
|
||||
{
|
||||
"segment_id": str(id),
|
||||
}
|
||||
)
|
||||
t = Table("segments")
|
||||
q = (
|
||||
self.querybuilder()
|
||||
.from_(t)
|
||||
.where(t.id == ParameterValue(self.uuid_to_db(id)))
|
||||
.delete()
|
||||
)
|
||||
with self.tx() as cur:
|
||||
# no need for explicit del from metadata table because of ON DELETE CASCADE
|
||||
sql, params = get_sql(q, self.parameter_format())
|
||||
sql = sql + " RETURNING id"
|
||||
result = cur.execute(sql, params).fetchone()
|
||||
if not result:
|
||||
raise NotFoundError(f"Segment {id} not found")
|
||||
|
||||
# Used by delete_collection to delete all segments for a collection along with
|
||||
# the collection itself in a single transaction.
|
||||
def delete_segments_for_collection(self, cur: Cursor, collection: UUID) -> None:
|
||||
segments_t = Table("segments")
|
||||
q = (
|
||||
self.querybuilder()
|
||||
.from_(segments_t)
|
||||
.where(segments_t.collection == ParameterValue(self.uuid_to_db(collection)))
|
||||
.delete()
|
||||
)
|
||||
sql, params = get_sql(q, self.parameter_format())
|
||||
cur.execute(sql, params)
|
||||
|
||||
@trace_method("SqlSysDB.delete_collection", OpenTelemetryGranularity.ALL)
|
||||
@override
|
||||
def delete_collection(
|
||||
self,
|
||||
id: UUID,
|
||||
tenant: str = DEFAULT_TENANT,
|
||||
database: str = DEFAULT_DATABASE,
|
||||
) -> None:
|
||||
"""Delete a collection and all associated segments from the SysDB. Deletes
|
||||
the log stream for this collection as well."""
|
||||
add_attributes_to_current_span(
|
||||
{
|
||||
"collection_id": str(id),
|
||||
}
|
||||
)
|
||||
t = Table("collections")
|
||||
databases_t = Table("databases")
|
||||
q = (
|
||||
self.querybuilder()
|
||||
.from_(t)
|
||||
.where(t.id == ParameterValue(self.uuid_to_db(id)))
|
||||
.where(
|
||||
t.database_id
|
||||
== self.querybuilder()
|
||||
.select(databases_t.id)
|
||||
.from_(databases_t)
|
||||
.where(databases_t.name == ParameterValue(database))
|
||||
.where(databases_t.tenant_id == ParameterValue(tenant))
|
||||
)
|
||||
.delete()
|
||||
)
|
||||
with self.tx() as cur:
|
||||
# no need for explicit del from metadata table because of ON DELETE CASCADE
|
||||
sql, params = get_sql(q, self.parameter_format())
|
||||
sql = sql + " RETURNING id"
|
||||
result = cur.execute(sql, params).fetchone()
|
||||
if not result:
|
||||
raise NotFoundError(f"Collection {id} not found")
|
||||
# Delete segments.
|
||||
self.delete_segments_for_collection(cur, id)
|
||||
|
||||
self._producer.delete_log(result[0])
|
||||
|
||||
@trace_method("SqlSysDB.update_segment", OpenTelemetryGranularity.ALL)
|
||||
@override
|
||||
def update_segment(
|
||||
self,
|
||||
collection: UUID,
|
||||
id: UUID,
|
||||
metadata: OptionalArgument[Optional[UpdateMetadata]] = Unspecified(),
|
||||
) -> None:
|
||||
add_attributes_to_current_span(
|
||||
{
|
||||
"segment_id": str(id),
|
||||
"collection": str(collection),
|
||||
}
|
||||
)
|
||||
segments_t = Table("segments")
|
||||
metadata_t = Table("segment_metadata")
|
||||
|
||||
q = (
|
||||
self.querybuilder()
|
||||
.update(segments_t)
|
||||
.where(segments_t.id == ParameterValue(self.uuid_to_db(id)))
|
||||
.set(segments_t.collection, ParameterValue(self.uuid_to_db(collection)))
|
||||
)
|
||||
|
||||
with self.tx() as cur:
|
||||
sql, params = get_sql(q, self.parameter_format())
|
||||
if sql: # pypika emits a blank string if nothing to do
|
||||
cur.execute(sql, params)
|
||||
|
||||
if metadata is None:
|
||||
q = (
|
||||
self.querybuilder()
|
||||
.from_(metadata_t)
|
||||
.where(metadata_t.segment_id == ParameterValue(self.uuid_to_db(id)))
|
||||
.delete()
|
||||
)
|
||||
sql, params = get_sql(q, self.parameter_format())
|
||||
cur.execute(sql, params)
|
||||
elif metadata != Unspecified():
|
||||
metadata = cast(UpdateMetadata, metadata)
|
||||
metadata = cast(UpdateMetadata, metadata)
|
||||
self._insert_metadata(
|
||||
cur,
|
||||
metadata_t,
|
||||
metadata_t.segment_id,
|
||||
id,
|
||||
metadata,
|
||||
set(metadata.keys()),
|
||||
)
|
||||
|
||||
@trace_method("SqlSysDB.update_collection", OpenTelemetryGranularity.ALL)
|
||||
@override
|
||||
def update_collection(
|
||||
self,
|
||||
id: UUID,
|
||||
name: OptionalArgument[str] = Unspecified(),
|
||||
dimension: OptionalArgument[Optional[int]] = Unspecified(),
|
||||
metadata: OptionalArgument[Optional[UpdateMetadata]] = Unspecified(),
|
||||
configuration: OptionalArgument[
|
||||
Optional[UpdateCollectionConfiguration]
|
||||
] = Unspecified(),
|
||||
) -> None:
|
||||
add_attributes_to_current_span(
|
||||
{
|
||||
"collection_id": str(id),
|
||||
}
|
||||
)
|
||||
collections_t = Table("collections")
|
||||
metadata_t = Table("collection_metadata")
|
||||
|
||||
q = (
|
||||
self.querybuilder()
|
||||
.update(collections_t)
|
||||
.where(collections_t.id == ParameterValue(self.uuid_to_db(id)))
|
||||
)
|
||||
|
||||
if not name == Unspecified():
|
||||
q = q.set(collections_t.name, ParameterValue(name))
|
||||
|
||||
if not dimension == Unspecified():
|
||||
q = q.set(collections_t.dimension, ParameterValue(dimension))
|
||||
|
||||
with self.tx() as cur:
|
||||
sql, params = get_sql(q, self.parameter_format())
|
||||
if sql: # pypika emits a blank string if nothing to do
|
||||
sql = sql + " RETURNING id"
|
||||
result = cur.execute(sql, params)
|
||||
if not result.fetchone():
|
||||
raise NotFoundError(f"Collection {id} not found")
|
||||
|
||||
# TODO: Update to use better semantics where it's possible to update
|
||||
# individual keys without wiping all the existing metadata.
|
||||
|
||||
# For now, follow current legancy semantics where metadata is fully reset
|
||||
if metadata != Unspecified():
|
||||
q = (
|
||||
self.querybuilder()
|
||||
.from_(metadata_t)
|
||||
.where(
|
||||
metadata_t.collection_id == ParameterValue(self.uuid_to_db(id))
|
||||
)
|
||||
.delete()
|
||||
)
|
||||
sql, params = get_sql(q, self.parameter_format())
|
||||
cur.execute(sql, params)
|
||||
if metadata is not None:
|
||||
metadata = cast(UpdateMetadata, metadata)
|
||||
self._insert_metadata(
|
||||
cur,
|
||||
metadata_t,
|
||||
metadata_t.collection_id,
|
||||
id,
|
||||
metadata,
|
||||
set(metadata.keys()),
|
||||
)
|
||||
|
||||
if configuration != Unspecified():
|
||||
update_configuration = cast(
|
||||
UpdateCollectionConfiguration, configuration
|
||||
)
|
||||
self._update_config_json_str(cur, update_configuration, id)
|
||||
else:
|
||||
if metadata != Unspecified():
|
||||
metadata = cast(UpdateMetadata, metadata)
|
||||
if metadata is not None:
|
||||
update_configuration = (
|
||||
update_collection_configuration_from_legacy_update_metadata(
|
||||
metadata
|
||||
)
|
||||
)
|
||||
self._update_config_json_str(cur, update_configuration, id)
|
||||
|
||||
def _update_config_json_str(
|
||||
self, cur: Cursor, update_configuration: UpdateCollectionConfiguration, id: UUID
|
||||
) -> None:
|
||||
collections_t = Table("collections")
|
||||
q = (
|
||||
self.querybuilder()
|
||||
.from_(collections_t)
|
||||
.select(collections_t.config_json_str)
|
||||
.where(collections_t.id == ParameterValue(self.uuid_to_db(id)))
|
||||
)
|
||||
sql, params = get_sql(q, self.parameter_format())
|
||||
row = cur.execute(sql, params).fetchone()
|
||||
if not row:
|
||||
raise NotFoundError(f"Collection {id} not found")
|
||||
config_json_str = row[0]
|
||||
existing_config = load_collection_configuration_from_json_str(config_json_str)
|
||||
new_config = overwrite_collection_configuration(
|
||||
existing_config, update_configuration
|
||||
)
|
||||
q = (
|
||||
self.querybuilder()
|
||||
.update(collections_t)
|
||||
.set(
|
||||
collections_t.config_json_str,
|
||||
ParameterValue(collection_configuration_to_json_str(new_config)),
|
||||
)
|
||||
.where(collections_t.id == ParameterValue(self.uuid_to_db(id)))
|
||||
)
|
||||
sql, params = get_sql(q, self.parameter_format())
|
||||
cur.execute(sql, params)
|
||||
|
||||
@trace_method("SqlSysDB._metadata_from_rows", OpenTelemetryGranularity.ALL)
|
||||
def _metadata_from_rows(
|
||||
self, rows: Sequence[Tuple[Any, ...]]
|
||||
) -> Optional[Metadata]:
|
||||
"""Given SQL rows, return a metadata map (assuming that the last four columns
|
||||
are the key, str_value, int_value & float_value)"""
|
||||
add_attributes_to_current_span(
|
||||
{
|
||||
"num_rows": len(rows),
|
||||
}
|
||||
)
|
||||
metadata: Dict[str, Union[str, int, float, bool]] = {}
|
||||
for row in rows:
|
||||
key = str(row[-5])
|
||||
if row[-4] is not None:
|
||||
metadata[key] = str(row[-4])
|
||||
elif row[-3] is not None:
|
||||
metadata[key] = int(row[-3])
|
||||
elif row[-2] is not None:
|
||||
metadata[key] = float(row[-2])
|
||||
elif row[-1] is not None:
|
||||
metadata[key] = bool(row[-1])
|
||||
return metadata or None
|
||||
|
||||
@trace_method("SqlSysDB._insert_metadata", OpenTelemetryGranularity.ALL)
|
||||
def _insert_metadata(
|
||||
self,
|
||||
cur: Cursor,
|
||||
table: Table,
|
||||
id_col: Column,
|
||||
id: UUID,
|
||||
metadata: UpdateMetadata,
|
||||
clear_keys: Optional[Set[str]] = None,
|
||||
) -> None:
|
||||
# It would be cleaner to use something like ON CONFLICT UPDATE here But that is
|
||||
# very difficult to do in a portable way (e.g sqlite and postgres have
|
||||
# completely different sytnax)
|
||||
add_attributes_to_current_span(
|
||||
{
|
||||
"num_keys": len(metadata),
|
||||
}
|
||||
)
|
||||
if clear_keys:
|
||||
q = (
|
||||
self.querybuilder()
|
||||
.from_(table)
|
||||
.where(id_col == ParameterValue(self.uuid_to_db(id)))
|
||||
.where(table.key.isin([ParameterValue(k) for k in clear_keys]))
|
||||
.delete()
|
||||
)
|
||||
sql, params = get_sql(q, self.parameter_format())
|
||||
cur.execute(sql, params)
|
||||
|
||||
q = (
|
||||
self.querybuilder()
|
||||
.into(table)
|
||||
.columns(
|
||||
id_col,
|
||||
table.key,
|
||||
table.str_value,
|
||||
table.int_value,
|
||||
table.float_value,
|
||||
table.bool_value,
|
||||
)
|
||||
)
|
||||
sql_id = self.uuid_to_db(id)
|
||||
for k, v in metadata.items():
|
||||
# Note: The order is important here because isinstance(v, bool)
|
||||
# and isinstance(v, int) both are true for v of bool type.
|
||||
if isinstance(v, bool):
|
||||
q = q.insert(
|
||||
ParameterValue(sql_id),
|
||||
ParameterValue(k),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
ParameterValue(int(v)),
|
||||
)
|
||||
elif isinstance(v, str):
|
||||
q = q.insert(
|
||||
ParameterValue(sql_id),
|
||||
ParameterValue(k),
|
||||
ParameterValue(v),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
)
|
||||
elif isinstance(v, int):
|
||||
q = q.insert(
|
||||
ParameterValue(sql_id),
|
||||
ParameterValue(k),
|
||||
None,
|
||||
ParameterValue(v),
|
||||
None,
|
||||
None,
|
||||
)
|
||||
elif isinstance(v, float):
|
||||
q = q.insert(
|
||||
ParameterValue(sql_id),
|
||||
ParameterValue(k),
|
||||
None,
|
||||
None,
|
||||
ParameterValue(v),
|
||||
None,
|
||||
)
|
||||
elif v is None:
|
||||
continue
|
||||
|
||||
sql, params = get_sql(q, self.parameter_format())
|
||||
if sql:
|
||||
cur.execute(sql, params)
|
||||
|
||||
def _insert_config_from_legacy_params(
|
||||
self, collection_id: Any, metadata: Optional[Metadata]
|
||||
) -> CollectionConfiguration:
|
||||
"""Insert the configuration from legacy metadata params into the collections table, and return the configuration object."""
|
||||
|
||||
# This is a legacy case where we don't have configuration stored in the database
|
||||
# This is non-destructive, we don't delete or overwrite any keys in the metadata
|
||||
|
||||
collections_t = Table("collections")
|
||||
|
||||
create_collection_config = CreateCollectionConfiguration()
|
||||
# Write the configuration into the database
|
||||
configuration_json_str = create_collection_configuration_to_json_str(
|
||||
create_collection_config, cast(CollectionMetadata, metadata)
|
||||
)
|
||||
q = (
|
||||
self.querybuilder()
|
||||
.update(collections_t)
|
||||
.set(
|
||||
collections_t.config_json_str,
|
||||
ParameterValue(configuration_json_str),
|
||||
)
|
||||
.where(collections_t.id == ParameterValue(collection_id))
|
||||
)
|
||||
sql, params = get_sql(q, self.parameter_format())
|
||||
with self.tx() as cur:
|
||||
cur.execute(sql, params)
|
||||
return load_collection_configuration_from_json_str(configuration_json_str)
|
||||
|
||||
@override
|
||||
def get_collection_size(self, id: UUID) -> int:
|
||||
raise NotImplementedError
|
||||
|
||||
@override
|
||||
def count_collections(
|
||||
self,
|
||||
tenant: str = DEFAULT_TENANT,
|
||||
database: Optional[str] = None,
|
||||
) -> int:
|
||||
"""Gets the number of collections for the (tenant, database) combination."""
|
||||
# TODO(Sanket): Implement this efficiently using a count query.
|
||||
# Note, the underlying get_collections api always requires a database
|
||||
# to be specified. In the sysdb implementation in go code, it does not
|
||||
# filter on database if it is set to "". This is a bad API and
|
||||
# should be fixed. For now, we will replicate the behavior.
|
||||
request_database: str = "" if database is None or database == "" else database
|
||||
return len(self.get_collections(tenant=tenant, database=request_database))
|
||||
Reference in New Issue
Block a user