mirror of
https://github.com/infiniflow/ragflow.git
synced 2026-08-19 23:05:10 +08:00
170 lines
6.0 KiB
Python
170 lines
6.0 KiB
Python
#
|
|
# Copyright 2024 The InfiniFlow Authors. All Rights Reserved.
|
|
#
|
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
|
# you may not use this file except in compliance with the License.
|
|
# You may obtain a copy of the License at
|
|
#
|
|
# http://www.apache.org/licenses/LICENSE-2.0
|
|
#
|
|
# Unless required by applicable law or agreed to in writing, software
|
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
# See the License for the specific language governing permissions and
|
|
# limitations under the License.
|
|
#
|
|
import operator
|
|
from functools import reduce
|
|
|
|
from playhouse.pool import PooledMySQLDatabase
|
|
|
|
from common.time_utils import current_timestamp, timestamp_to_date
|
|
|
|
from api.db.db_models import DB, DataBaseModel, is_gaussdb_compatible_database
|
|
|
|
|
|
def _gaussdb_replace_insert_by_id(model, batch, preserve):
|
|
"""Replace Peewee ``ON CONFLICT`` upserts on A/ORA-compatible GaussDB.
|
|
|
|
Peewee's PostgreSQL dialect renders
|
|
``insert_many(...).on_conflict(conflict_target="id")`` as
|
|
``INSERT ... ON CONFLICT (...) DO UPDATE ... RETURNING``. A/ORA-compatible
|
|
GaussDB rejects that syntax, which breaks APIs such as
|
|
``/datasets/{id}/documents/parse`` when they batch-write tasks.
|
|
|
|
For ``DB_TYPE=gaussdb`` only, split this replace-by-id operation into two
|
|
steps:
|
|
|
|
* Query the IDs already present in the batch and update those rows using
|
|
the same columns that Peewee's PostgreSQL path would preserve.
|
|
* Insert rows with new IDs without generating unsupported ``ON CONFLICT``
|
|
SQL.
|
|
|
|
MySQL retains its on-duplicate-key behavior, and PostgreSQL/OceanBase retain
|
|
Peewee's existing generation path. SQL for other databases is unchanged.
|
|
"""
|
|
ids = [data.get("id") for data in batch if data.get("id") is not None]
|
|
existing_ids = set()
|
|
if ids:
|
|
existing_ids = {row[0] for row in model.select(model.id).where(model.id.in_(ids)).tuples()}
|
|
|
|
insert_rows = []
|
|
update_columns = [column for column in preserve if column != "id"]
|
|
for data in batch:
|
|
row_id = data.get("id")
|
|
if row_id in existing_ids:
|
|
update_payload = {column: data[column] for column in update_columns if column in data}
|
|
if update_payload:
|
|
model.update(update_payload).where(model.id == row_id).execute()
|
|
else:
|
|
insert_rows.append(data)
|
|
|
|
if insert_rows:
|
|
model.insert_many(insert_rows).execute()
|
|
|
|
|
|
@DB.connection_context()
|
|
def bulk_insert_into_db(model, data_source, replace_on_conflict=False):
|
|
DB.create_tables([model])
|
|
|
|
for i, data in enumerate(data_source):
|
|
current_time = current_timestamp() + i
|
|
current_date = timestamp_to_date(current_time)
|
|
if "create_time" not in data:
|
|
data["create_time"] = current_time
|
|
data["create_date"] = timestamp_to_date(data["create_time"])
|
|
data["update_time"] = current_time
|
|
data["update_date"] = current_date
|
|
|
|
preserve = tuple(data_source[0].keys() - {"create_time", "create_date"})
|
|
|
|
batch_size = 1000
|
|
|
|
for i in range(0, len(data_source), batch_size):
|
|
with DB.atomic():
|
|
query = model.insert_many(data_source[i : i + batch_size])
|
|
if replace_on_conflict:
|
|
if isinstance(DB, PooledMySQLDatabase):
|
|
query = query.on_conflict(preserve=preserve)
|
|
elif is_gaussdb_compatible_database():
|
|
# Peewee's PostgreSQL path emits `ON CONFLICT (...) DO
|
|
# UPDATE ... RETURNING`, which A/ORA-compatible GaussDB
|
|
# rejects. Query existing IDs first, then UPDATE or INSERT
|
|
# to preserve the operation without leaking PostgreSQL SQL.
|
|
_gaussdb_replace_insert_by_id(model, data_source[i : i + batch_size], preserve)
|
|
continue
|
|
else:
|
|
query = query.on_conflict(conflict_target="id", preserve=preserve)
|
|
query.execute()
|
|
|
|
|
|
def get_dynamic_db_model(base, job_id):
|
|
return type(base.model(table_index=get_dynamic_tracking_table_index(job_id=job_id)))
|
|
|
|
|
|
def get_dynamic_tracking_table_index(job_id):
|
|
return job_id[:8]
|
|
|
|
|
|
def fill_db_model_object(model_object, human_model_dict):
|
|
for k, v in human_model_dict.items():
|
|
attr_name = "f_%s" % k
|
|
if hasattr(model_object.__class__, attr_name):
|
|
setattr(model_object, attr_name, v)
|
|
return model_object
|
|
|
|
|
|
# https://docs.peewee-orm.com/en/latest/peewee/query_operators.html
|
|
supported_operators = {
|
|
"==": operator.eq,
|
|
"<": operator.lt,
|
|
"<=": operator.le,
|
|
">": operator.gt,
|
|
">=": operator.ge,
|
|
"!=": operator.ne,
|
|
"<<": operator.lshift,
|
|
">>": operator.rshift,
|
|
"%": operator.mod,
|
|
"**": operator.pow,
|
|
"^": operator.xor,
|
|
"~": operator.inv,
|
|
}
|
|
|
|
|
|
def query_dict2expression(model: type[DataBaseModel], query: dict[str, bool | int | str | list | tuple]):
|
|
expression = []
|
|
|
|
for field, value in query.items():
|
|
if not isinstance(value, (list, tuple)):
|
|
value = ("==", value)
|
|
op, *val = value
|
|
|
|
field = getattr(model, f"f_{field}")
|
|
value = supported_operators[op](field, val[0]) if op in supported_operators else getattr(field, op)(*val)
|
|
expression.append(value)
|
|
|
|
return reduce(operator.iand, expression)
|
|
|
|
|
|
def query_db(model: type[DataBaseModel], limit: int = 0, offset: int = 0, query: dict = None, order_by: str | list | tuple | None = None):
|
|
data = model.select()
|
|
if query:
|
|
data = data.where(query_dict2expression(model, query))
|
|
count = data.count()
|
|
|
|
if not order_by:
|
|
order_by = "create_time"
|
|
if not isinstance(order_by, (list, tuple)):
|
|
order_by = (order_by, "asc")
|
|
order_by, order = order_by
|
|
order_by = getattr(model, f"f_{order_by}")
|
|
order_by = getattr(order_by, order)()
|
|
data = data.order_by(order_by)
|
|
|
|
if limit > 0:
|
|
data = data.limit(limit)
|
|
if offset > 0:
|
|
data = data.offset(offset)
|
|
|
|
return list(data), count
|