refac/enh: db session sharing

This commit is contained in:
Timothy Jaeryang Baek
2025-12-29 00:21:18 +04:00
parent 6dd0f99b90
commit b1d0f00d8c
23 changed files with 1173 additions and 663 deletions
@@ -30,27 +30,27 @@ from sqlalchemy.exc import NoSuchTableError
from sqlalchemy.dialects.postgresql.psycopg2 import PGDialect_psycopg2
from sqlalchemy.dialects import registry
class OpenGaussDialect(PGDialect_psycopg2):
name = "opengauss"
def _get_server_version_info(self, connection):
try:
version = connection.exec_driver_sql("SELECT version()").scalar()
if not version:
return (9, 0, 0)
match = re.search(
r"openGauss\s+(\d+)\.(\d+)\.(\d+)(?:-\w+)?",
version,
re.IGNORECASE
r"openGauss\s+(\d+)\.(\d+)\.(\d+)(?:-\w+)?", version, re.IGNORECASE
)
if match:
return (int(match.group(1)), int(match.group(2)), int(match.group(3)))
return super()._get_server_version_info(connection)
except Exception:
return (9, 0, 0)
# Register dialect
registry.register("opengauss", __name__, "OpenGaussDialect")
@@ -78,6 +78,7 @@ Base = declarative_base()
log = logging.getLogger(__name__)
log.setLevel(SRC_LOG_LEVELS["RAG"])
class DocumentChunk(Base):
__tablename__ = "document_chunk"
@@ -87,29 +88,30 @@ class DocumentChunk(Base):
text = Column(Text, nullable=True)
vmetadata = Column(MutableDict.as_mutable(JSONB), nullable=True)
class OpenGaussClient(VectorDBBase):
def __init__(self) -> None:
if not OPENGAUSS_DB_URL:
from open_webui.internal.db import Session
self.session = Session
from open_webui.internal.db import ScopedSession
self.session = ScopedSession
else:
engine_kwargs = {
"pool_pre_ping": True,
"dialect": OpenGaussDialect()
}
engine_kwargs = {"pool_pre_ping": True, "dialect": OpenGaussDialect()}
if isinstance(OPENGAUSS_POOL_SIZE, int) and OPENGAUSS_POOL_SIZE > 0:
engine_kwargs.update({
"pool_size": OPENGAUSS_POOL_SIZE,
"max_overflow": OPENGAUSS_POOL_MAX_OVERFLOW,
"pool_timeout": OPENGAUSS_POOL_TIMEOUT,
"pool_recycle": OPENGAUSS_POOL_RECYCLE,
"poolclass": QueuePool
})
engine_kwargs.update(
{
"pool_size": OPENGAUSS_POOL_SIZE,
"max_overflow": OPENGAUSS_POOL_MAX_OVERFLOW,
"pool_timeout": OPENGAUSS_POOL_TIMEOUT,
"pool_recycle": OPENGAUSS_POOL_RECYCLE,
"poolclass": QueuePool,
}
)
else:
engine_kwargs["poolclass"] = NullPool
engine = create_engine(OPENGAUSS_DB_URL,** engine_kwargs)
engine = create_engine(OPENGAUSS_DB_URL, **engine_kwargs)
SessionLocal = sessionmaker(
autocommit=False, autoflush=False, bind=engine, expire_on_commit=False
@@ -160,7 +162,9 @@ class OpenGaussClient(VectorDBBase):
else:
raise Exception("The 'vector' column type is not Vector.")
else:
raise Exception("The 'vector' column does not exist in the 'document_chunk' table.")
raise Exception(
"The 'vector' column does not exist in the 'document_chunk' table."
)
def adjust_vector_length(self, vector: List[float]) -> List[float]:
current_length = len(vector)
@@ -185,7 +189,9 @@ class OpenGaussClient(VectorDBBase):
new_items.append(new_chunk)
self.session.bulk_save_objects(new_items)
self.session.commit()
log.info(f"Inserting {len(new_items)} items into collection '{collection_name}'.")
log.info(
f"Inserting {len(new_items)} items into collection '{collection_name}'."
)
except Exception as e:
self.session.rollback()
log.exception(f"Failed to insert data: {e}")
@@ -215,7 +221,9 @@ class OpenGaussClient(VectorDBBase):
)
self.session.add(new_chunk)
self.session.commit()
log.info(f"Inserting/updating {len(items)} items in collection '{collection_name}'.")
log.info(
f"Inserting/updating {len(items)} items in collection '{collection_name}'."
)
except Exception as e:
self.session.rollback()
log.exception(f"Failed to insert or update data.: {e}")
@@ -241,7 +249,9 @@ class OpenGaussClient(VectorDBBase):
q_vector_col = column("q_vector", Vector(VECTOR_LENGTH))
query_vectors = (
values(qid_col, q_vector_col)
.data([(idx, vector_expr(vector)) for idx, vector in enumerate(vectors)])
.data(
[(idx, vector_expr(vector)) for idx, vector in enumerate(vectors)]
)
.alias("query_vectors")
)
@@ -249,13 +259,17 @@ class OpenGaussClient(VectorDBBase):
DocumentChunk.id,
DocumentChunk.text,
DocumentChunk.vmetadata,
(DocumentChunk.vector.cosine_distance(query_vectors.c.q_vector)).label("distance"),
(DocumentChunk.vector.cosine_distance(query_vectors.c.q_vector)).label(
"distance"
),
]
subq = (
select(*result_fields)
.where(DocumentChunk.collection_name == collection_name)
.order_by(DocumentChunk.vector.cosine_distance(query_vectors.c.q_vector))
.order_by(
DocumentChunk.vector.cosine_distance(query_vectors.c.q_vector)
)
)
if limit is not None:
subq = subq.limit(limit)
@@ -368,7 +382,9 @@ class OpenGaussClient(VectorDBBase):
query = query.filter(DocumentChunk.id.in_(ids))
if filter:
for key, value in filter.items():
query = query.filter(DocumentChunk.vmetadata[key].astext == str(value))
query = query.filter(
DocumentChunk.vmetadata[key].astext == str(value)
)
deleted = query.delete(synchronize_session=False)
self.session.commit()
log.info(f"Deleted {deleted} items from collection '{collection_name}'")
@@ -395,7 +411,8 @@ class OpenGaussClient(VectorDBBase):
exists = (
self.session.query(DocumentChunk)
.filter(DocumentChunk.collection_name == collection_name)
.first() is not None
.first()
is not None
)
self.session.rollback()
return exists
@@ -406,4 +423,4 @@ class OpenGaussClient(VectorDBBase):
def delete_collection(self, collection_name: str) -> None:
self.delete(collection_name)
log.info(f"Collection '{collection_name}' has been deleted")
log.info(f"Collection '{collection_name}' has been deleted")
@@ -90,9 +90,9 @@ class PgvectorClient(VectorDBBase):
# if no pgvector uri, use the existing database connection
if not PGVECTOR_DB_URL:
from open_webui.internal.db import Session
from open_webui.internal.db import ScopedSession
self.session = Session
self.session = ScopedSession
else:
if isinstance(PGVECTOR_POOL_SIZE, int):
if PGVECTOR_POOL_SIZE > 0: