refac
This commit is contained in:
@@ -16,7 +16,7 @@ log = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class File(Base):
|
||||
__tablename__ = "file"
|
||||
__tablename__ = 'file'
|
||||
id = Column(String, primary_key=True, unique=True)
|
||||
user_id = Column(String)
|
||||
hash = Column(Text, nullable=True)
|
||||
@@ -58,9 +58,9 @@ class FileMeta(BaseModel):
|
||||
content_type: Optional[str] = None
|
||||
size: Optional[int] = None
|
||||
|
||||
model_config = ConfigDict(extra="allow")
|
||||
model_config = ConfigDict(extra='allow')
|
||||
|
||||
@model_validator(mode="before")
|
||||
@model_validator(mode='before')
|
||||
@classmethod
|
||||
def sanitize_meta(cls, data):
|
||||
"""Sanitize metadata fields to handle malformed legacy data."""
|
||||
@@ -68,14 +68,12 @@ class FileMeta(BaseModel):
|
||||
return data
|
||||
|
||||
# Handle content_type that may be a list like ['application/pdf', None]
|
||||
content_type = data.get("content_type")
|
||||
content_type = data.get('content_type')
|
||||
if isinstance(content_type, list):
|
||||
# Extract first non-None string value
|
||||
data["content_type"] = next(
|
||||
(item for item in content_type if isinstance(item, str)), None
|
||||
)
|
||||
data['content_type'] = next((item for item in content_type if isinstance(item, str)), None)
|
||||
elif content_type is not None and not isinstance(content_type, str):
|
||||
data["content_type"] = None
|
||||
data['content_type'] = None
|
||||
|
||||
return data
|
||||
|
||||
@@ -92,7 +90,7 @@ class FileModelResponse(BaseModel):
|
||||
created_at: int # timestamp in epoch
|
||||
updated_at: Optional[int] = None # timestamp in epoch, optional for legacy files
|
||||
|
||||
model_config = ConfigDict(extra="allow")
|
||||
model_config = ConfigDict(extra='allow')
|
||||
|
||||
|
||||
class FileMetadataResponse(BaseModel):
|
||||
@@ -123,25 +121,22 @@ class FileUpdateForm(BaseModel):
|
||||
meta: Optional[dict] = None
|
||||
|
||||
|
||||
|
||||
class FilesTable:
|
||||
def insert_new_file(
|
||||
self, user_id: str, form_data: FileForm, db: Optional[Session] = None
|
||||
) -> Optional[FileModel]:
|
||||
def insert_new_file(self, user_id: str, form_data: FileForm, db: Optional[Session] = None) -> Optional[FileModel]:
|
||||
with get_db_context(db) as db:
|
||||
file_data = form_data.model_dump()
|
||||
|
||||
# Sanitize meta to remove non-JSON-serializable objects
|
||||
# (e.g. callable tool functions, MCP client instances from middleware)
|
||||
if file_data.get("meta"):
|
||||
file_data["meta"] = sanitize_metadata(file_data["meta"])
|
||||
if file_data.get('meta'):
|
||||
file_data['meta'] = sanitize_metadata(file_data['meta'])
|
||||
|
||||
file = FileModel(
|
||||
**{
|
||||
**file_data,
|
||||
"user_id": user_id,
|
||||
"created_at": int(time.time()),
|
||||
"updated_at": int(time.time()),
|
||||
'user_id': user_id,
|
||||
'created_at': int(time.time()),
|
||||
'updated_at': int(time.time()),
|
||||
}
|
||||
)
|
||||
|
||||
@@ -155,12 +150,10 @@ class FilesTable:
|
||||
else:
|
||||
return None
|
||||
except Exception as e:
|
||||
log.exception(f"Error inserting a new file: {e}")
|
||||
log.exception(f'Error inserting a new file: {e}')
|
||||
return None
|
||||
|
||||
def get_file_by_id(
|
||||
self, id: str, db: Optional[Session] = None
|
||||
) -> Optional[FileModel]:
|
||||
def get_file_by_id(self, id: str, db: Optional[Session] = None) -> Optional[FileModel]:
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
try:
|
||||
@@ -171,9 +164,7 @@ class FilesTable:
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
def get_file_by_id_and_user_id(
|
||||
self, id: str, user_id: str, db: Optional[Session] = None
|
||||
) -> Optional[FileModel]:
|
||||
def get_file_by_id_and_user_id(self, id: str, user_id: str, db: Optional[Session] = None) -> Optional[FileModel]:
|
||||
with get_db_context(db) as db:
|
||||
try:
|
||||
file = db.query(File).filter_by(id=id, user_id=user_id).first()
|
||||
@@ -184,9 +175,7 @@ class FilesTable:
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
def get_file_metadata_by_id(
|
||||
self, id: str, db: Optional[Session] = None
|
||||
) -> Optional[FileMetadataResponse]:
|
||||
def get_file_metadata_by_id(self, id: str, db: Optional[Session] = None) -> Optional[FileMetadataResponse]:
|
||||
with get_db_context(db) as db:
|
||||
try:
|
||||
file = db.get(File, id)
|
||||
@@ -204,9 +193,7 @@ class FilesTable:
|
||||
with get_db_context(db) as db:
|
||||
return [FileModel.model_validate(file) for file in db.query(File).all()]
|
||||
|
||||
def check_access_by_user_id(
|
||||
self, id, user_id, permission="write", db: Optional[Session] = None
|
||||
) -> bool:
|
||||
def check_access_by_user_id(self, id, user_id, permission='write', db: Optional[Session] = None) -> bool:
|
||||
file = self.get_file_by_id(id, db=db)
|
||||
if not file:
|
||||
return False
|
||||
@@ -215,21 +202,14 @@ class FilesTable:
|
||||
# Implement additional access control logic here as needed
|
||||
return False
|
||||
|
||||
def get_files_by_ids(
|
||||
self, ids: list[str], db: Optional[Session] = None
|
||||
) -> list[FileModel]:
|
||||
def get_files_by_ids(self, ids: list[str], db: Optional[Session] = None) -> list[FileModel]:
|
||||
with get_db_context(db) as db:
|
||||
return [
|
||||
FileModel.model_validate(file)
|
||||
for file in db.query(File)
|
||||
.filter(File.id.in_(ids))
|
||||
.order_by(File.updated_at.desc())
|
||||
.all()
|
||||
for file in db.query(File).filter(File.id.in_(ids)).order_by(File.updated_at.desc()).all()
|
||||
]
|
||||
|
||||
def get_file_metadatas_by_ids(
|
||||
self, ids: list[str], db: Optional[Session] = None
|
||||
) -> list[FileMetadataResponse]:
|
||||
def get_file_metadatas_by_ids(self, ids: list[str], db: Optional[Session] = None) -> list[FileMetadataResponse]:
|
||||
with get_db_context(db) as db:
|
||||
return [
|
||||
FileMetadataResponse(
|
||||
@@ -239,22 +219,15 @@ class FilesTable:
|
||||
created_at=file.created_at,
|
||||
updated_at=file.updated_at,
|
||||
)
|
||||
for file in db.query(
|
||||
File.id, File.hash, File.meta, File.created_at, File.updated_at
|
||||
)
|
||||
for file in db.query(File.id, File.hash, File.meta, File.created_at, File.updated_at)
|
||||
.filter(File.id.in_(ids))
|
||||
.order_by(File.updated_at.desc())
|
||||
.all()
|
||||
]
|
||||
|
||||
def get_files_by_user_id(
|
||||
self, user_id: str, db: Optional[Session] = None
|
||||
) -> list[FileModel]:
|
||||
def get_files_by_user_id(self, user_id: str, db: Optional[Session] = None) -> list[FileModel]:
|
||||
with get_db_context(db) as db:
|
||||
return [
|
||||
FileModel.model_validate(file)
|
||||
for file in db.query(File).filter_by(user_id=user_id).all()
|
||||
]
|
||||
return [FileModel.model_validate(file) for file in db.query(File).filter_by(user_id=user_id).all()]
|
||||
|
||||
def get_file_list(
|
||||
self,
|
||||
@@ -262,7 +235,7 @@ class FilesTable:
|
||||
skip: int = 0,
|
||||
limit: int = 50,
|
||||
db: Optional[Session] = None,
|
||||
) -> "FileListResponse":
|
||||
) -> 'FileListResponse':
|
||||
with get_db_context(db) as db:
|
||||
query = db.query(File)
|
||||
if user_id:
|
||||
@@ -272,10 +245,7 @@ class FilesTable:
|
||||
|
||||
items = [
|
||||
FileModel.model_validate(file)
|
||||
for file in query.order_by(File.updated_at.desc(), File.id.desc())
|
||||
.offset(skip)
|
||||
.limit(limit)
|
||||
.all()
|
||||
for file in query.order_by(File.updated_at.desc(), File.id.desc()).offset(skip).limit(limit).all()
|
||||
]
|
||||
|
||||
return FileListResponse(items=items, total=total)
|
||||
@@ -296,17 +266,17 @@ class FilesTable:
|
||||
A SQL LIKE compatible pattern with proper escaping.
|
||||
"""
|
||||
# Escape SQL special characters first, then convert glob wildcards
|
||||
pattern = glob.replace("\\", "\\\\")
|
||||
pattern = pattern.replace("%", "\\%")
|
||||
pattern = pattern.replace("_", "\\_")
|
||||
pattern = pattern.replace("*", "%")
|
||||
pattern = pattern.replace("?", "_")
|
||||
pattern = glob.replace('\\', '\\\\')
|
||||
pattern = pattern.replace('%', '\\%')
|
||||
pattern = pattern.replace('_', '\\_')
|
||||
pattern = pattern.replace('*', '%')
|
||||
pattern = pattern.replace('?', '_')
|
||||
return pattern
|
||||
|
||||
def search_files(
|
||||
self,
|
||||
user_id: Optional[str] = None,
|
||||
filename: str = "*",
|
||||
filename: str = '*',
|
||||
skip: int = 0,
|
||||
limit: int = 100,
|
||||
db: Optional[Session] = None,
|
||||
@@ -331,15 +301,12 @@ class FilesTable:
|
||||
query = query.filter_by(user_id=user_id)
|
||||
|
||||
pattern = self._glob_to_like_pattern(filename)
|
||||
if pattern != "%":
|
||||
query = query.filter(File.filename.ilike(pattern, escape="\\"))
|
||||
if pattern != '%':
|
||||
query = query.filter(File.filename.ilike(pattern, escape='\\'))
|
||||
|
||||
return [
|
||||
FileModel.model_validate(file)
|
||||
for file in query.order_by(File.created_at.desc(), File.id.desc())
|
||||
.offset(skip)
|
||||
.limit(limit)
|
||||
.all()
|
||||
for file in query.order_by(File.created_at.desc(), File.id.desc()).offset(skip).limit(limit).all()
|
||||
]
|
||||
|
||||
def update_file_by_id(
|
||||
@@ -362,12 +329,10 @@ class FilesTable:
|
||||
db.commit()
|
||||
return FileModel.model_validate(file)
|
||||
except Exception as e:
|
||||
log.exception(f"Error updating file completely by id: {e}")
|
||||
log.exception(f'Error updating file completely by id: {e}')
|
||||
return None
|
||||
|
||||
def update_file_hash_by_id(
|
||||
self, id: str, hash: Optional[str], db: Optional[Session] = None
|
||||
) -> Optional[FileModel]:
|
||||
def update_file_hash_by_id(self, id: str, hash: Optional[str], db: Optional[Session] = None) -> Optional[FileModel]:
|
||||
with get_db_context(db) as db:
|
||||
try:
|
||||
file = db.query(File).filter_by(id=id).first()
|
||||
@@ -379,9 +344,7 @@ class FilesTable:
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
def update_file_data_by_id(
|
||||
self, id: str, data: dict, db: Optional[Session] = None
|
||||
) -> Optional[FileModel]:
|
||||
def update_file_data_by_id(self, id: str, data: dict, db: Optional[Session] = None) -> Optional[FileModel]:
|
||||
with get_db_context(db) as db:
|
||||
try:
|
||||
file = db.query(File).filter_by(id=id).first()
|
||||
@@ -390,12 +353,9 @@ class FilesTable:
|
||||
db.commit()
|
||||
return FileModel.model_validate(file)
|
||||
except Exception as e:
|
||||
|
||||
return None
|
||||
|
||||
def update_file_metadata_by_id(
|
||||
self, id: str, meta: dict, db: Optional[Session] = None
|
||||
) -> Optional[FileModel]:
|
||||
def update_file_metadata_by_id(self, id: str, meta: dict, db: Optional[Session] = None) -> Optional[FileModel]:
|
||||
with get_db_context(db) as db:
|
||||
try:
|
||||
file = db.query(File).filter_by(id=id).first()
|
||||
|
||||
Reference in New Issue
Block a user