mini-rag / src /models /ChunkModel.py
mustaphaelkady's picture
Clean deploy mini-rag to Hugging Face
338036b
Raw
History Blame Contribute Delete
2.63 kB
from .BaseDataModel import BaseDataModel
from .db_schemes.minirag.schemes import DataChunk
from sqlalchemy.future import select
from sqlalchemy import func, delete
from bson.objectid import ObjectId
class ChunkModel(BaseDataModel):
def __init__(self, db_client: object):
super().__init__(db_client=db_client)
self.db_client = db_client
@classmethod
async def create_instance(cls, db_client: object):
instance = cls(db_client)
return instance
async def create_chunk(self, chunk: DataChunk):
async with self.db_client() as session:
async with session.begin():
session.add(chunk)
await session.commit()
await session.refresh(chunk)
return chunk
async def get_chunk(self, chunk_id: str):
async with self.db_client() as session:
result = await session.execute(select(DataChunk.chunk_id).where(DataChunk.chunk_id==chunk_id))
chunk = result.scalar_one_or_none()
return chunk
async def insert_many_chunks(self, chunks: list, batch_size: int=100):
async with self.db_client() as session:
async with session.begin():
for i in range(0, len(chunks), batch_size):
chunk = chunks[i:i+batch_size]
session.add_all(chunk)
await session.commit()
return len(chunks)
async def delete_chunks_by_project_id(self, project_id: ObjectId):
async with self.db_client() as session:
stmt = (delete(DataChunk).where(DataChunk.chunk_project_id==project_id))
result = await session.execute(stmt)
await session.commit()
return result.rowcount
async def get_poject_chunks(self, project_id: ObjectId, page_no: int=1, page_size: int=20):
async with self.db_client() as session:
stmt = select(DataChunk).where(DataChunk.chunk_project_id==project_id).order_by(DataChunk.chunk_order).offset((page_no-1)*page_size).limit(page_size)
result = await session.execute(stmt)
records = result.scalars().all()
return records
async def get_chunk_count(self,project_id:ObjectId):
total_count = 0
async with self.db_client() as session:
count_sql = select(func.count(DataChunk.chunk_id)).where(DataChunk.chunk_project_id==project_id)
result = await session.execute(count_sql)
total_count = result.scalar()
return total_count