from typing import Optional, List, Tuple, Any, Dict from sqlalchemy import select, func, update as sa_update, desc, asc, text, case, literal_column, cast, String, or_, and_ from sqlalchemy.ext.asyncio import AsyncSession from app.base_service import BaseService from app.db_compat import get_db_type, json_extract, json_has_key from zq_smart_table.model import SmartTable, SmartField, SmartRecord, SmartView, SmartTableLink, SmartTableComment, SmartDocumentVersion, SmartDocumentTemplate, SmartWikiSpace from zq_smart_table.schema import ( SmartTableCreate, SmartTableUpdate, SmartFieldCreate, SmartFieldUpdate, SmartRecordCreate, SmartRecordUpdate, SmartViewCreate, SmartViewUpdate, RecordFilterRule, RecordSortRule, LinkedRecordItem, WikiSpaceCreate, WikiSpaceUpdate, ) class SmartTableService(BaseService[SmartTable, SmartTableCreate, SmartTableUpdate]): model = SmartTable RESOURCE_TYPE = "smart_table" @classmethod async def get_user_tables( cls, db: AsyncSession, page: int = 1, page_size: int = 1000, wiki_space_id: Optional[str] = None, user_id: Optional[str] = None, dept_id: Optional[str] = None, role_ids: Optional[List[str]] = None, is_superuser: bool = False, ) -> Tuple[List[SmartTable], int]: from zq_smart_table.permission.model import SmartTableCollaborator base_query = select(SmartTable).where(SmartTable.is_deleted == False) # noqa: E712 if wiki_space_id is not None: base_query = base_query.where(SmartTable.wiki_space_id == wiki_space_id) else: base_query = base_query.where(SmartTable.wiki_space_id.is_(None)) if not is_superuser and user_id: subject_conds = [ and_( SmartTableCollaborator.subject_type == "user", SmartTableCollaborator.subject_id == user_id, ) ] if dept_id: subject_conds.append( and_( SmartTableCollaborator.subject_type == "dept", SmartTableCollaborator.subject_id == dept_id, ) ) if role_ids: for rid in role_ids: subject_conds.append( and_( SmartTableCollaborator.subject_type == "role", SmartTableCollaborator.subject_id == rid, ) ) collab_table_ids = ( select(SmartTableCollaborator.table_id) .where( SmartTableCollaborator.is_deleted == False, # noqa: E712 or_(*subject_conds), ) .distinct() ) base_query = base_query.where( or_( SmartTable.sys_creator_id == user_id, SmartTable.id.in_(collab_table_ids), ) ) count_result = await db.execute( select(func.count()).select_from(base_query.subquery()) ) total = count_result.scalar() or 0 offset = (page - 1) * page_size result = await db.execute( base_query.order_by(desc(SmartTable.sort), desc(SmartTable.sys_create_datetime)) .offset(offset) .limit(page_size) ) items = list(result.scalars().all()) return items, total class SmartFieldService(BaseService[SmartField, SmartFieldCreate, SmartFieldUpdate]): model = SmartField RESOURCE_TYPE = "smart_field" @classmethod async def get_by_table(cls, db: AsyncSession, table_id: str) -> List[SmartField]: result = await db.execute( select(SmartField) .where(SmartField.table_id == table_id, SmartField.is_deleted == False) # noqa: E712 .order_by(SmartField.sort, SmartField.sys_create_datetime) ) return list(result.scalars().all()) @classmethod async def reorder(cls, db: AsyncSession, table_id: str, field_ids: List[str]) -> None: """批量排序:一条 CASE WHEN SQL 替代 N 次 UPDATE""" if not field_ids: return whens = [(SmartField.id == fid, idx) for idx, fid in enumerate(field_ids)] await db.execute( sa_update(SmartField) .where(SmartField.table_id == table_id, SmartField.id.in_(field_ids)) .values(sort=case(*whens, else_=SmartField.sort)) ) await db.commit() @classmethod async def delete_and_clean(cls, db: AsyncSession, field_id: str) -> bool: """删除字段并用一条 SQL 清理所有 record 的 values""" field = await cls.get_by_id(db, field_id) if not field: return False db_type = get_db_type() if db_type == "postgresql": await db.execute( text( "UPDATE smart_record SET \"values\" = \"values\" - :field_id " "WHERE table_id = :table_id AND is_deleted = false " "AND \"values\" \\? :field_id" ), {"field_id": field_id, "table_id": field.table_id}, ) else: await db.execute( text( "UPDATE smart_record SET `values` = JSON_REMOVE(`values`, CONCAT('$.', :field_id)) " "WHERE table_id = :table_id AND is_deleted = 0 " "AND JSON_CONTAINS_PATH(`values`, 'one', CONCAT('$.', :field_id))" ), {"field_id": field_id, "table_id": field.table_id}, ) field.is_deleted = True await db.commit() return True @classmethod async def get_next_sort(cls, db: AsyncSession, table_id: str) -> int: result = await db.execute( select(func.coalesce(func.max(SmartField.sort), -1)) .where(SmartField.table_id == table_id, SmartField.is_deleted == False) # noqa: E712 ) return (result.scalar() or 0) + 1 class SmartRecordService(BaseService[SmartRecord, SmartRecordCreate, SmartRecordUpdate]): model = SmartRecord RESOURCE_TYPE = "smart_record" @classmethod async def get_by_table_cursor( cls, db: AsyncSession, table_id: str, cursor: Optional[str] = None, limit: int = 200, ) -> Tuple[List[SmartRecord], Optional[str], int]: """ 游标分页查询(适用于大数据量滚动加载)。 返回 (records, next_cursor, total) """ base = select(SmartRecord).where( SmartRecord.table_id == table_id, SmartRecord.is_deleted == False, # noqa: E712 ) count_result = await db.execute( select(func.count()).select_from(base.subquery()) ) total = count_result.scalar() or 0 query = base.order_by(SmartRecord.sort, SmartRecord.id) if cursor: query = query.where(SmartRecord.id > cursor) query = query.limit(limit + 1) result = await db.execute(query) items = list(result.scalars().all()) next_cursor = None if len(items) > limit: items = items[:limit] next_cursor = items[-1].id return items, next_cursor, total @classmethod async def get_by_table( cls, db: AsyncSession, table_id: str, page: int = 1, page_size: int = 200 ) -> Tuple[List[SmartRecord], int]: filters = [SmartRecord.table_id == table_id] return await cls.get_list(db, page=page, page_size=page_size, filters=filters) @classmethod async def update_cell( cls, db: AsyncSession, record_id: str, field_id: str, value: Any ) -> Optional[SmartRecord]: """使用数据库原生 JSON 原子更新,避免读-改-写""" import json db_type = get_db_type() json_value = json.dumps(value) if db_type == "postgresql": await db.execute( text( "UPDATE smart_record SET \"values\" = jsonb_set(" "COALESCE(\"values\", CAST('{}' AS jsonb)), " ":path, CAST(:val AS jsonb)), " "sys_update_datetime = now() " "WHERE id = :id AND is_deleted = false" ), {"path": [field_id], "val": json_value, "id": record_id}, ) else: await db.execute( text( "UPDATE smart_record SET `values` = JSON_SET(" "COALESCE(`values`, '{}'), " "CONCAT('$.', :field_id), CAST(:val AS JSON)), " "sys_update_datetime = NOW() " "WHERE id = :id AND is_deleted = 0" ), {"field_id": field_id, "val": json_value, "id": record_id}, ) await db.commit() record = await cls.get_by_id(db, record_id) return record @classmethod async def batch_update_cells( cls, db: AsyncSession, record_id: str, cells: Dict[str, Any] ) -> Optional[SmartRecord]: """批量更新多个单元格(一次 SQL)""" import json record = await cls.get_by_id(db, record_id) if not record: return None current_values = dict(record.values) if record.values else {} current_values.update(cells) record.values = current_values await db.commit() await db.refresh(record) return record @classmethod async def batch_update_multi_records( cls, db: AsyncSession, updates: List[Dict[str, Any]] ) -> int: """ 批量更新多条记录的单元格,合并为单次事务。 updates: [{ "record_id": "...", "cells": { fieldId: value, ... } }, ...] 返回成功更新的记录数。 """ import json as _json db_type = get_db_type() count = 0 for item in updates: record_id = item.get("record_id") cells = item.get("cells", {}) if not record_id or not cells: continue if db_type == "postgresql": merge_obj = _json.dumps(cells) await db.execute( text( 'UPDATE smart_record SET "values" = ' 'COALESCE("values", \'{}\'::jsonb) || CAST(:merge AS jsonb), ' 'sys_update_datetime = now() ' 'WHERE id = :id AND is_deleted = false' ), {"merge": merge_obj, "id": record_id}, ) else: record = await cls.get_by_id(db, record_id) if not record: continue current = dict(record.values) if record.values else {} current.update(cells) record.values = current count += 1 await db.commit() return count @classmethod async def batch_remove_field(cls, db: AsyncSession, table_id: str, field_id: str) -> None: """用一条 SQL 从所有 record 中移除某个 field 的值""" db_type = get_db_type() if db_type == "postgresql": await db.execute( text( "UPDATE smart_record SET \"values\" = \"values\" - :field_id " "WHERE table_id = :table_id AND is_deleted = false " "AND \"values\" \\? :field_id" ), {"field_id": field_id, "table_id": table_id}, ) else: await db.execute( text( "UPDATE smart_record SET `values` = JSON_REMOVE(`values`, CONCAT('$.', :field_id)) " "WHERE table_id = :table_id AND is_deleted = 0" ), {"field_id": field_id, "table_id": table_id}, ) await db.commit() @classmethod async def get_deleted_records( cls, db: AsyncSession, table_id: str, page: int = 1, page_size: int = 50, ) -> Tuple[List[SmartRecord], int]: """查询已删除的记录(回收站)""" base = select(SmartRecord).where( SmartRecord.table_id == table_id, SmartRecord.is_deleted == True, # noqa: E712 ) count_result = await db.execute( select(func.count()).select_from(base.subquery()) ) total = count_result.scalar() or 0 query = base.order_by(desc(SmartRecord.sys_update_datetime)).offset( (page - 1) * page_size ).limit(page_size) result = await db.execute(query) return list(result.scalars().all()), total @classmethod async def batch_restore_records( cls, db: AsyncSession, table_id: str, record_ids: List[str] ) -> int: """批量恢复已删除的记录(仅限指定表)""" result = await db.execute( sa_update(SmartRecord) .where( SmartRecord.id.in_(record_ids), SmartRecord.table_id == table_id, SmartRecord.is_deleted == True, # noqa: E712 ) .values(is_deleted=False) ) await db.commit() return result.rowcount # type: ignore @classmethod async def permanent_delete(cls, db: AsyncSession, table_id: str, record_id: str) -> bool: """永久删除记录(仅限指定表的已删除记录)""" from sqlalchemy import delete as sa_del result = await db.execute( sa_del(SmartRecord).where( SmartRecord.id == record_id, SmartRecord.table_id == table_id, SmartRecord.is_deleted == True, # noqa: E712 ) ) await db.commit() return result.rowcount > 0 # type: ignore @classmethod async def empty_trash(cls, db: AsyncSession, table_id: str) -> int: """清空回收站""" from sqlalchemy import delete as sa_del result = await db.execute( sa_del(SmartRecord).where( SmartRecord.table_id == table_id, SmartRecord.is_deleted == True, # noqa: E712 ) ) await db.commit() return result.rowcount # type: ignore @classmethod async def get_next_auto_number(cls, db: AsyncSession, table_id: str, field_id: str) -> int: """用 SQL 聚合查询获取 AutoNumber 最大值,避免全表加载到 Python""" db_type = get_db_type() if db_type == "postgresql": result = await db.execute( text( "SELECT COALESCE(MAX(CAST(\"values\" ->> :field_id AS int)), 0) " "FROM smart_record " "WHERE table_id = :table_id AND is_deleted = false " "AND \"values\" \\? :field_id " "AND (\"values\" ->> :field_id) ~ '^[0-9]+$'" ), {"field_id": field_id, "table_id": table_id}, ) else: result = await db.execute( text( "SELECT COALESCE(MAX(CAST(JSON_UNQUOTE(JSON_EXTRACT(`values`, CONCAT('$.', :field_id))) AS SIGNED)), 0) " "FROM smart_record " "WHERE table_id = :table_id AND is_deleted = 0 " "AND JSON_CONTAINS_PATH(`values`, 'one', CONCAT('$.', :field_id))" ), {"field_id": field_id, "table_id": table_id}, ) return (result.scalar() or 0) + 1 @classmethod async def get_record_count(cls, db: AsyncSession, table_id: str) -> int: result = await db.execute( select(func.count()).where( SmartRecord.table_id == table_id, SmartRecord.is_deleted == False, # noqa: E712 ) ) return result.scalar() or 0 # ==================== Server-side filter / sort / search ==================== @classmethod def _build_filter_condition(cls, rule: RecordFilterRule): """将单条 RecordFilterRule 转换为 SQLAlchemy 条件表达式""" col = json_extract(SmartRecord.values, rule.field_id) op = rule.operator val = rule.value if op == "isEmpty": return or_( ~json_has_key(SmartRecord.values, rule.field_id), col == None, # noqa: E711 col == "", ) if op == "isNotEmpty": return and_( json_has_key(SmartRecord.values, rule.field_id), col != None, # noqa: E711 col != "", ) if op == "equals": return col == str(val) if val is not None else col == None # noqa: E711 if op == "notEquals": return col != str(val) if val is not None else col != None # noqa: E711 if op == "contains": return col.ilike(f"%{val}%") if val else col == col if op == "notContains": return ~col.ilike(f"%{val}%") if val else col == col if op == "greaterThan": return cast(col, String) > str(val) if op == "lessThan": return cast(col, String) < str(val) if op == "greaterThanOrEqual": return cast(col, String) >= str(val) if op == "lessThanOrEqual": return cast(col, String) <= str(val) return True # noqa: fallback @classmethod def _build_search_conditions(cls, keyword: str, search_field_ids: Optional[List[str]] = None): """构建全文搜索条件:在指定字段或所有字段中匹配关键词""" if not keyword: return None like_pattern = f"%{keyword}%" if search_field_ids: conds = [json_extract(SmartRecord.values, fid).ilike(like_pattern) for fid in search_field_ids] else: # 对 JSONB 整列做 cast(text) ILIKE,兼容所有字段 conds = [cast(SmartRecord.values, String).ilike(like_pattern)] return or_(*conds) if conds else None @classmethod async def get_by_table_cursor_filtered( cls, db: AsyncSession, table_id: str, filters: Optional[List[RecordFilterRule]] = None, filter_logic: str = "and", sorts: Optional[List[RecordSortRule]] = None, search: Optional[str] = None, search_field_ids: Optional[List[str]] = None, extra_conditions: Optional[list] = None, cursor: Optional[str] = None, limit: int = 200, skip_count: bool = False, ) -> Tuple[List[SmartRecord], Optional[str], int]: """ 服务端筛选 + 排序 + 搜索 + 游标分页。 extra_conditions: 外部传入的额外 SQLAlchemy 条件(如行权限) skip_count: 为True时跳过count查询(用于加载更多时提升性能) """ base = select(SmartRecord).where( SmartRecord.table_id == table_id, SmartRecord.is_deleted == False, # noqa: E712 ) has_complex_filter = False if filters: filter_conds = [cls._build_filter_condition(f) for f in filters] if filter_logic == "or": base = base.where(or_(*filter_conds)) else: base = base.where(and_(*filter_conds)) has_complex_filter = True if search and search.strip(): search_cond = cls._build_search_conditions(search.strip(), search_field_ids) if search_cond is not None: base = base.where(search_cond) has_complex_filter = True if extra_conditions: for cond in extra_conditions: base = base.where(cond) has_complex_filter = True total = -1 if not skip_count: if has_complex_filter: count_result = await db.execute( select(func.count()).select_from(base.subquery()) ) else: count_result = await db.execute( select(func.count()).where( SmartRecord.table_id == table_id, SmartRecord.is_deleted == False, # noqa: E712 ) ) total = count_result.scalar() or 0 query = base if sorts: for s in sorts: col = json_extract(SmartRecord.values, s.field_id) query = query.order_by(desc(col) if s.direction == "desc" else asc(col)) query = query.order_by(SmartRecord.sort, SmartRecord.id) if cursor: query = query.where(SmartRecord.id > cursor) query = query.limit(limit + 1) result = await db.execute(query) items = list(result.scalars().all()) next_cursor = None if len(items) > limit: items = items[:limit] next_cursor = items[-1].id return items, next_cursor, total @classmethod async def reorder(cls, db: AsyncSession, table_id: str, record_ids: List[str]) -> None: """批量更新记录排序""" if not record_ids: return whens = [(SmartRecord.id == rid, idx) for idx, rid in enumerate(record_ids)] await db.execute( sa_update(SmartRecord) .where(SmartRecord.table_id == table_id, SmartRecord.id.in_(record_ids)) .values(sort=case(*whens, else_=SmartRecord.sort)) ) await db.commit() class SmartViewService(BaseService[SmartView, SmartViewCreate, SmartViewUpdate]): model = SmartView RESOURCE_TYPE = "smart_view" @classmethod async def get_by_table(cls, db: AsyncSession, table_id: str) -> List[SmartView]: result = await db.execute( select(SmartView) .where(SmartView.table_id == table_id, SmartView.is_deleted == False) # noqa: E712 .order_by(SmartView.sort, SmartView.sys_create_datetime) ) return list(result.scalars().all()) # ==================== Import / Export ==================== class SmartExportService: """CSV / Excel 导出""" @classmethod def _format_cell(cls, val: Any, field_type: str) -> str: if val is None: return "" if field_type == "link" and isinstance(val, list): return ", ".join( (item.get("title", "") if isinstance(item, dict) else str(item)) for item in val ) if isinstance(val, list): return ", ".join(str(v) for v in val) return str(val) @classmethod async def export_csv( cls, db: AsyncSession, table_id: str, fields: List[SmartField], records: List[SmartRecord], ) -> str: import csv import io output = io.StringIO() writer = csv.writer(output) writer.writerow([f.name for f in fields]) for r in records: row = [] for f in fields: val = (r.values or {}).get(f.id, "") row.append(cls._format_cell(val, f.type)) writer.writerow(row) return output.getvalue() @classmethod async def export_xlsx( cls, db: AsyncSession, table_id: str, fields: List[SmartField], records: List[SmartRecord], ) -> bytes: from openpyxl import Workbook import io wb = Workbook() ws = wb.active ws.title = "Sheet1" ws.append([f.name for f in fields]) for r in records: row = [] for f in fields: val = (r.values or {}).get(f.id, "") row.append(cls._format_cell(val, f.type)) ws.append(row) buf = io.BytesIO() wb.save(buf) return buf.getvalue() class SmartImportService: """CSV / Excel 导入""" BATCH_SIZE = 500 @classmethod async def import_csv( cls, db: AsyncSession, table_id: str, file_content: bytes, fields: List[SmartField], ) -> Dict[str, Any]: import csv import io text = file_content.decode("utf-8-sig") reader = csv.reader(io.StringIO(text)) headers = next(reader, None) if not headers: return {"success": 0, "fail": 0, "errors": ["文件为空或无表头"]} return await cls._import_rows(db, table_id, headers, list(reader), fields) @classmethod async def import_xlsx( cls, db: AsyncSession, table_id: str, file_content: bytes, fields: List[SmartField], ) -> Dict[str, Any]: from openpyxl import load_workbook import io wb = load_workbook(io.BytesIO(file_content), read_only=True) ws = wb.active rows_iter = ws.iter_rows(values_only=True) header_row = next(rows_iter, None) if not header_row: return {"success": 0, "fail": 0, "errors": ["文件为空或无表头"]} headers = [str(h) if h else "" for h in header_row] data_rows = [[str(c) if c is not None else "" for c in row] for row in rows_iter] return await cls._import_rows(db, table_id, headers, data_rows, fields) @classmethod async def _import_rows( cls, db: AsyncSession, table_id: str, headers: List[str], rows: List[List[str]], fields: List[SmartField], ) -> Dict[str, Any]: field_map: Dict[str, SmartField] = {f.name: f for f in fields} col_to_field: List[Optional[SmartField]] = [] new_fields_created = 0 for h in headers: h = h.strip() if h in field_map: col_to_field.append(field_map[h]) else: new_field = SmartField( table_id=table_id, name=h, type="text", width=150, visible=True, required=False, config={}, sort=len(fields) + new_fields_created, ) db.add(new_field) await db.flush() field_map[h] = new_field col_to_field.append(new_field) new_fields_created += 1 success = 0 errors: List[str] = [] batch: List[SmartRecord] = [] for row_idx, row in enumerate(rows): try: values: Dict[str, Any] = {} for col_idx, cell_val in enumerate(row): if col_idx < len(col_to_field) and col_to_field[col_idx]: field = col_to_field[col_idx] if field.type in ("link", "lookup", "rollup"): continue values[field.id] = cls._parse_cell(cell_val, field.type) record = SmartRecord(table_id=table_id, values=values) batch.append(record) success += 1 if len(batch) >= cls.BATCH_SIZE: db.add_all(batch) await db.commit() batch = [] except Exception as e: errors.append(f"第{row_idx + 2}行: {str(e)}") if batch: db.add_all(batch) await db.commit() return { "success": success, "fail": len(errors), "errors": errors[:50], "new_fields": new_fields_created, } @classmethod def _parse_cell(cls, val: str, field_type: str) -> Any: if not val: return None if field_type == "number": try: return float(val) if "." in val else int(val) except ValueError: return val if field_type == "checkbox": return val.lower() in ("true", "1", "yes", "是") if field_type == "rating": try: return int(float(val)) except ValueError: return 0 return val # ==================== Link / Lookup / Rollup ==================== class SmartLinkService: """管理 Link / Lookup / Rollup 字段的创建、删除、关联值解析""" # ---------- Link 字段创建 / 删除 ---------- @classmethod async def create_link_field( cls, db: AsyncSession, table_id: str, field_name: str, target_table_id: str, sort: int = 0, ) -> Tuple[SmartField, SmartField]: """ 创建 Link 字段及目标表的对称 Link 字段。 返回 (source_field, symmetric_field)。 """ source_table = await SmartTableService.get_by_id(db, table_id) target_table = await SmartTableService.get_by_id(db, target_table_id) if not source_table or not target_table: raise ValueError("源表或目标表不存在") source_field_data = SmartFieldCreate( table_id=table_id, name=field_name, type="link", sort=sort or await SmartFieldService.get_next_sort(db, table_id), config={ "linkedTableId": target_table_id, }, ) source_field = await SmartFieldService.create(db, source_field_data) sym_field_data = SmartFieldCreate( table_id=target_table_id, name=source_table.name, type="link", sort=await SmartFieldService.get_next_sort(db, target_table_id), config={ "linkedTableId": table_id, "symmetricFieldId": source_field.id, }, ) sym_field = await SmartFieldService.create(db, sym_field_data) source_field.config = { **(source_field.config or {}), "symmetricFieldId": sym_field.id, } await db.commit() await db.refresh(source_field) return source_field, sym_field @classmethod async def delete_link_field(cls, db: AsyncSession, field_id: str) -> bool: """删除 Link 字段、对称字段和所有关联记录""" field = await SmartFieldService.get_by_id(db, field_id) if not field or field.type != "link": return False sym_field_id = (field.config or {}).get("symmetricFieldId") await db.execute( sa_update(SmartTableLink) .where(SmartTableLink.field_id == field_id) .values(is_deleted=True) ) if sym_field_id: await db.execute( sa_update(SmartTableLink) .where(SmartTableLink.field_id == sym_field_id) .values(is_deleted=True) ) sym_field = await SmartFieldService.get_by_id(db, sym_field_id) if sym_field: sym_field.is_deleted = True field.is_deleted = True await db.commit() return True # ---------- 关联记录操作 ---------- @classmethod async def set_linked_records( cls, db: AsyncSession, field_id: str, source_record_id: str, target_record_ids: List[str], ) -> None: """ 全量设置关联:先移除旧关联,再插入新关联。 同时维护对称方向。 """ field = await SmartFieldService.get_by_id(db, field_id) if not field or field.type != "link": return sym_field_id = (field.config or {}).get("symmetricFieldId") existing = await db.execute( select(SmartTableLink).where( SmartTableLink.field_id == field_id, SmartTableLink.source_record_id == source_record_id, SmartTableLink.is_deleted == False, # noqa: E712 ) ) existing_links = list(existing.scalars().all()) existing_target_ids = {l.target_record_id for l in existing_links} to_add = set(target_record_ids) - existing_target_ids to_remove = existing_target_ids - set(target_record_ids) for link in existing_links: if link.target_record_id in to_remove: link.is_deleted = True if sym_field_id and to_remove: sym_result = await db.execute( select(SmartTableLink).where( SmartTableLink.field_id == sym_field_id, SmartTableLink.target_record_id == source_record_id, SmartTableLink.source_record_id.in_(to_remove), SmartTableLink.is_deleted == False, # noqa: E712 ) ) for sym_link in sym_result.scalars().all(): sym_link.is_deleted = True for tid in to_add: db.add(SmartTableLink( field_id=field_id, source_record_id=source_record_id, target_record_id=tid, )) if sym_field_id: db.add(SmartTableLink( field_id=sym_field_id, source_record_id=tid, target_record_id=source_record_id, )) await db.commit() @classmethod async def get_linked_record_ids( cls, db: AsyncSession, field_id: str, source_record_id: str ) -> List[str]: result = await db.execute( select(SmartTableLink.target_record_id).where( SmartTableLink.field_id == field_id, SmartTableLink.source_record_id == source_record_id, SmartTableLink.is_deleted == False, # noqa: E712 ) ) return list(result.scalars().all()) # ---------- 解析关联值(批量注入到 records) ---------- @classmethod async def resolve_link_values( cls, db: AsyncSession, fields: List[SmartField], records: List[SmartRecord], ) -> None: """批量解析 Link 字段值,注入 [{id, title}] 到 record.values""" link_fields = [f for f in fields if f.type == "link"] if not link_fields or not records: return record_ids = [r.id for r in records] for lf in link_fields: result = await db.execute( select(SmartTableLink.source_record_id, SmartTableLink.target_record_id) .where( SmartTableLink.field_id == lf.id, SmartTableLink.source_record_id.in_(record_ids), SmartTableLink.is_deleted == False, # noqa: E712 ) ) links = result.all() target_ids = list({row[1] for row in links}) title_map: Dict[str, str] = {} if target_ids: target_table_id = (lf.config or {}).get("linkedTableId") if target_table_id: title_map = await cls._get_record_titles(db, target_table_id, target_ids) src_map: Dict[str, List[LinkedRecordItem]] = {} for src_id, tgt_id in links: src_map.setdefault(src_id, []).append( LinkedRecordItem(id=tgt_id, title=title_map.get(tgt_id, "")) ) for r in records: vals = dict(r.values) if r.values else {} items = src_map.get(r.id, []) vals[lf.id] = [item.model_dump() for item in items] r.values = vals @classmethod async def resolve_lookup_values( cls, db: AsyncSession, fields: List[SmartField], records: List[SmartRecord], ) -> None: """解析 Lookup 字段值:通过 linkFieldId 找到关联记录,取 lookupFieldId 的值""" lookup_fields = [f for f in fields if f.type == "lookup"] if not lookup_fields or not records: return record_ids = [r.id for r in records] for lkf in lookup_fields: cfg = lkf.config or {} link_field_id = cfg.get("linkFieldId") lookup_field_id = cfg.get("lookupFieldId") if not link_field_id or not lookup_field_id: continue result = await db.execute( select(SmartTableLink.source_record_id, SmartTableLink.target_record_id) .where( SmartTableLink.field_id == link_field_id, SmartTableLink.source_record_id.in_(record_ids), SmartTableLink.is_deleted == False, # noqa: E712 ) ) links = result.all() target_ids = list({row[1] for row in links}) target_values: Dict[str, Any] = {} if target_ids: tgt_result = await db.execute( select(SmartRecord.id, SmartRecord.values).where( SmartRecord.id.in_(target_ids), SmartRecord.is_deleted == False, # noqa: E712 ) ) for rid, vals in tgt_result.all(): target_values[rid] = (vals or {}).get(lookup_field_id) src_map: Dict[str, list] = {} for src_id, tgt_id in links: v = target_values.get(tgt_id) if v is not None: src_map.setdefault(src_id, []).append(v) for r in records: vals = dict(r.values) if r.values else {} vals[lkf.id] = src_map.get(r.id, []) r.values = vals @classmethod async def resolve_rollup_values( cls, db: AsyncSession, fields: List[SmartField], records: List[SmartRecord], ) -> None: """解析 Rollup 字段值:聚合关联记录的某个字段""" rollup_fields = [f for f in fields if f.type == "rollup"] if not rollup_fields or not records: return record_ids = [r.id for r in records] for rf in rollup_fields: cfg = rf.config or {} link_field_id = cfg.get("linkFieldId") rollup_field_id = cfg.get("rollupFieldId") aggregation = cfg.get("aggregation", "COUNT") if not link_field_id or not rollup_field_id: continue result = await db.execute( select(SmartTableLink.source_record_id, SmartTableLink.target_record_id) .where( SmartTableLink.field_id == link_field_id, SmartTableLink.source_record_id.in_(record_ids), SmartTableLink.is_deleted == False, # noqa: E712 ) ) links = result.all() target_ids = list({row[1] for row in links}) target_values: Dict[str, Any] = {} if target_ids: tgt_result = await db.execute( select(SmartRecord.id, SmartRecord.values).where( SmartRecord.id.in_(target_ids), SmartRecord.is_deleted == False, # noqa: E712 ) ) for rid, vals in tgt_result.all(): target_values[rid] = (vals or {}).get(rollup_field_id) src_groups: Dict[str, list] = {} for src_id, tgt_id in links: v = target_values.get(tgt_id) src_groups.setdefault(src_id, []).append(v) for r in records: vals = dict(r.values) if r.values else {} group = src_groups.get(r.id, []) vals[rf.id] = cls._aggregate(group, aggregation) r.values = vals @classmethod def _aggregate(cls, values: List[Any], aggregation: str) -> Any: agg = aggregation.upper() if agg == "COUNT": return len(values) if agg == "COUNTA": return len([v for v in values if v is not None and v != ""]) nums = [] for v in values: if v is None: continue try: nums.append(float(v)) except (TypeError, ValueError): pass if agg == "SUM": return sum(nums) if nums else 0 if agg == "AVG": return sum(nums) / len(nums) if nums else None if agg == "MIN": return min(nums) if nums else None if agg == "MAX": return max(nums) if nums else None return len(values) # ---------- 记录删除时清理关联 ---------- @classmethod async def cleanup_record_links(cls, db: AsyncSession, record_id: str) -> None: """删除记录时,软删除所有涉及该记录的关联""" await db.execute( sa_update(SmartTableLink) .where( or_( SmartTableLink.source_record_id == record_id, SmartTableLink.target_record_id == record_id, ), SmartTableLink.is_deleted == False, # noqa: E712 ) .values(is_deleted=True) ) # ---------- 搜索目标表记录 ---------- @classmethod async def search_records( cls, db: AsyncSession, table_id: str, keyword: str = "", limit: int = 20, ) -> List[Dict[str, Any]]: """搜索指定表的记录,返回 [{id, title}]""" fields = await SmartFieldService.get_by_table(db, table_id) primary_field = fields[0] if fields else None base = select(SmartRecord).where( SmartRecord.table_id == table_id, SmartRecord.is_deleted == False, # noqa: E712 ) if keyword.strip() and primary_field: col = json_extract(SmartRecord.values, primary_field.id) base = base.where(col.ilike(f"%{keyword.strip()}%")) base = base.order_by(SmartRecord.sort, SmartRecord.id).limit(limit) result = await db.execute(base) rows = list(result.scalars().all()) items = [] for r in rows: title = "" if primary_field: title = str((r.values or {}).get(primary_field.id, "") or "") items.append({"id": r.id, "title": title}) return items # ---------- 内部工具 ---------- @classmethod async def _get_record_titles( cls, db: AsyncSession, table_id: str, record_ids: List[str] ) -> Dict[str, str]: """批量获取记录的标题(取第一个字段的值)""" fields = await SmartFieldService.get_by_table(db, table_id) primary_field = fields[0] if fields else None if not primary_field: return {} result = await db.execute( select(SmartRecord.id, SmartRecord.values).where( SmartRecord.id.in_(record_ids), SmartRecord.is_deleted == False, # noqa: E712 ) ) title_map: Dict[str, str] = {} for rid, vals in result.all(): title_map[rid] = str((vals or {}).get(primary_field.id, "") or "") return title_map class SmartSummaryService: """字段汇总聚合计算 - 合并多字段聚合为尽量少的SQL""" SUPPORTED = {"SUM", "AVG", "MIN", "MAX", "COUNT", "COUNTA", "COUNT_EMPTY", "PERCENT_EMPTY", "PERCENT_FILLED"} @classmethod async def compute( cls, db: AsyncSession, table_id: str, field_agg_map: Dict[str, str], filters: Optional[List[RecordFilterRule]] = None, filter_logic: str = "and", search: Optional[str] = None, extra_conditions: Optional[list] = None, ) -> Tuple[Dict[str, Any], int]: """ 对指定字段执行聚合计算,将同类聚合合并为单条SQL减少查询次数。 field_agg_map: {fieldId: aggregation_type} 返回 ({fieldId: result_value}, total_count) """ base = select(SmartRecord).where( SmartRecord.table_id == table_id, SmartRecord.is_deleted == False, # noqa: E712 ) if filters: filter_conds = [SmartRecordService._build_filter_condition(f) for f in filters] if filter_logic == "or": base = base.where(or_(*filter_conds)) else: base = base.where(and_(*filter_conds)) if search and search.strip(): search_cond = SmartRecordService._build_search_conditions(search.strip()) if search_cond is not None: base = base.where(search_cond) if extra_conditions: for cond in extra_conditions: base = base.where(cond) subq = base.subquery() count_result = await db.execute(select(func.count()).select_from(subq)) total = count_result.scalar() or 0 if total == 0: return {fid: None for fid in field_agg_map}, 0 summaries: Dict[str, Any] = {} db_type = get_db_type() numeric_fields: Dict[str, str] = {} count_needed_fields: Dict[str, str] = {} for field_id, agg_type in field_agg_map.items(): agg_upper = agg_type.upper() if agg_upper not in cls.SUPPORTED: summaries[field_id] = None elif agg_upper == "COUNT": summaries[field_id] = total elif agg_upper in ("SUM", "AVG", "MIN", "MAX"): numeric_fields[field_id] = agg_upper else: count_needed_fields[field_id] = agg_upper alias = subq.alias("sub") agg_fn_map = {"SUM": func.sum, "AVG": func.avg, "MIN": func.min, "MAX": func.max} all_columns = [] all_labels: List[Tuple[str, str, str]] = [] for field_id, agg_type in numeric_fields.items(): if db_type == "postgresql": num_expr = text(f"CAST(sub.\"values\" ->> '{field_id}' AS NUMERIC)") cond = text(f"(sub.\"values\" ->> '{field_id}') ~ '^-?[0-9]+(\\.[0-9]+)?$'") else: num_expr = text( f"CAST(JSON_UNQUOTE(JSON_EXTRACT(sub.`values`, '$.{field_id}')) AS DECIMAL(20,6))" ) cond = text( f"JSON_UNQUOTE(JSON_EXTRACT(sub.`values`, '$.{field_id}')) REGEXP '^-?[0-9]+(\\\\.[0-9]+)?$'" ) agg_fn = agg_fn_map[agg_type] label = f"num_{field_id}" all_columns.append(agg_fn(case((cond, num_expr), else_=None)).label(label)) all_labels.append((field_id, label, "numeric")) unique_count_fields = set(count_needed_fields.keys()) for field_id in unique_count_fields: if db_type == "postgresql": not_null_cond = text( f"sub.\"values\" ->> '{field_id}' IS NOT NULL " f"AND sub.\"values\" ->> '{field_id}' != ''" ) else: not_null_cond = text( f"JSON_EXTRACT(sub.`values`, '$.{field_id}') IS NOT NULL " f"AND JSON_UNQUOTE(JSON_EXTRACT(sub.`values`, '$.{field_id}')) != ''" ) label = f"cnt_{field_id}" all_columns.append(func.count(case((not_null_cond, 1), else_=None)).label(label)) all_labels.append((field_id, label, "count")) if all_columns: stmt = select(*all_columns).select_from(alias) row = (await db.execute(stmt)).first() if row: for field_id, label, col_type in all_labels: val = getattr(row, label, None) if col_type == "numeric": summaries[field_id] = round(float(val), 6) if val is not None else None else: non_empty = val or 0 agg_type = count_needed_fields[field_id] if agg_type == "COUNTA": summaries[field_id] = non_empty elif agg_type == "COUNT_EMPTY": summaries[field_id] = total - non_empty elif agg_type == "PERCENT_EMPTY": summaries[field_id] = round((total - non_empty) / total * 100, 1) if total > 0 else 0 elif agg_type == "PERCENT_FILLED": summaries[field_id] = round(non_empty / total * 100, 1) if total > 0 else 0 return summaries, total class SmartFormulaService: """公式字段计算服务:在返回记录时动态计算公式字段的值""" @classmethod def resolve_formula_values( cls, fields: List[SmartField], records: list ) -> None: """遍历公式字段,对每条记录计算公式值并注入 record.values(同步,不需要 DB)""" from zq_smart_table.formula import compute_formula formula_fields = [ f for f in fields if f.type == "formula" and f.config and f.config.get("formula") ] if not formula_fields: return field_name_map = {f.name: f.id for f in fields} for record in records: values = record.values if record.values else {} for ff in formula_fields: formula_str = ff.config["formula"] result = compute_formula(formula_str, values, field_name_map) result_type = ff.config.get("formulaResultType", "text") values[ff.id] = cls._cast_result(result, result_type) record.values = values @staticmethod def _cast_result(value: Any, result_type: str) -> Any: if value == "#ERROR": return value if result_type == "number": try: return float(value) if value is not None else None except (ValueError, TypeError): return "#ERROR" if result_type == "boolean": if isinstance(value, bool): return value return bool(value) if value is not None else False if result_type == "date": return str(value) if value is not None else None return str(value) if value is not None else "" class SmartValidationService: """字段数据校验服务""" @classmethod async def validate_cell( cls, db: AsyncSession, field: SmartField, value: Any, record_id: Optional[str] = None, ) -> Optional[str]: """ 校验单个单元格值,返回错误消息或 None(通过)。 校验规则存储在 field.config["validation"] 中。 """ import re as _re config = field.config or {} validation = config.get("validation") if not validation: return None if field.required and (value is None or value == "" or value == []): custom = validation.get("message") return custom or f"「{field.name}」不能为空" if value is None or value == "" or value == []: return None if "min" in validation and validation["min"] is not None: try: num_val = float(value) if num_val < float(validation["min"]): return validation.get("message") or f"「{field.name}」不能小于 {validation['min']}" except (ValueError, TypeError): pass if "max" in validation and validation["max"] is not None: try: num_val = float(value) if num_val > float(validation["max"]): return validation.get("message") or f"「{field.name}」不能大于 {validation['max']}" except (ValueError, TypeError): pass if "minLength" in validation and validation["minLength"] is not None: str_val = str(value) if len(str_val) < int(validation["minLength"]): return validation.get("message") or f"「{field.name}」长度不能少于 {validation['minLength']} 个字符" if "maxLength" in validation and validation["maxLength"] is not None: str_val = str(value) if len(str_val) > int(validation["maxLength"]): return validation.get("message") or f"「{field.name}」长度不能超过 {validation['maxLength']} 个字符" if "pattern" in validation and validation["pattern"]: str_val = str(value) try: if not _re.fullmatch(validation["pattern"], str_val): return validation.get("message") or f"「{field.name}」格式不正确" except _re.error: pass if validation.get("unique"): is_dup = await cls._check_unique(db, field, value, record_id) if is_dup: return validation.get("message") or f"「{field.name}」的值已存在,不能重复" return None @classmethod async def _check_unique( cls, db: AsyncSession, field: SmartField, value: Any, exclude_record_id: Optional[str] = None, ) -> bool: db_type = get_db_type() str_val = str(value) if db_type == "postgresql": cond = text(f"\"values\" ->> '{field.id}' = :v") else: cond = text(f"JSON_UNQUOTE(JSON_EXTRACT(`values`, '$.{field.id}')) = :v") stmt = ( select(func.count()) .select_from(SmartRecord.__table__) .where( SmartRecord.table_id == field.table_id, SmartRecord.is_deleted == False, # noqa: E712 cond.bindparams(v=str_val), ) ) if exclude_record_id: stmt = stmt.where(SmartRecord.id != exclude_record_id) result = await db.execute(stmt) return (result.scalar() or 0) > 0 @classmethod async def validate_cells_batch( cls, db: AsyncSession, table_id: str, cells: Dict[str, Any], record_id: Optional[str] = None, ) -> Dict[str, str]: """批量校验多个单元格,返回 {fieldId: errorMessage}""" fields = await SmartFieldService.get_by_table(db, table_id) field_map = {f.id: f for f in fields} errors: Dict[str, str] = {} for field_id, value in cells.items(): field = field_map.get(field_id) if not field: continue err = await cls.validate_cell(db, field, value, record_id) if err: errors[field_id] = err return errors class SmartCommentService: """记录评论服务""" @classmethod async def get_by_record(cls, db: AsyncSession, record_id: str) -> List[SmartTableComment]: result = await db.execute( select(SmartTableComment) .where(SmartTableComment.record_id == record_id, SmartTableComment.is_deleted == False) # noqa: E712 .order_by(SmartTableComment.sys_create_datetime.asc()) ) return list(result.scalars().all()) @classmethod async def create(cls, db: AsyncSession, record_id: str, user_id: str, content: str, mentions: List[str] = None, parent_id: str = None) -> SmartTableComment: comment = SmartTableComment( record_id=record_id, user_id=user_id, content=content, mentions=mentions or [], parent_id=parent_id, sys_creator_id=user_id, ) db.add(comment) await db.commit() await db.refresh(comment) return comment @classmethod async def update(cls, db: AsyncSession, comment_id: str, user_id: str, content: str, mentions: List[str] = None) -> Optional[SmartTableComment]: result = await db.execute( select(SmartTableComment) .where(SmartTableComment.id == comment_id, SmartTableComment.is_deleted == False) # noqa: E712 ) comment = result.scalar_one_or_none() if not comment or comment.user_id != user_id: return None comment.content = content comment.mentions = mentions or [] comment.sys_modifier_id = user_id await db.commit() await db.refresh(comment) return comment @classmethod async def delete(cls, db: AsyncSession, comment_id: str, user_id: str) -> bool: result = await db.execute( select(SmartTableComment) .where(SmartTableComment.id == comment_id, SmartTableComment.is_deleted == False) # noqa: E712 ) comment = result.scalar_one_or_none() if not comment or comment.user_id != user_id: return False comment.is_deleted = True comment.sys_modifier_id = user_id await db.commit() return True class SmartDocumentVersionService: """文档版本历史服务""" MAX_VERSIONS_PER_DOC = 100 @classmethod async def get_next_version(cls, db: AsyncSession, document_id: str) -> int: result = await db.execute( select(func.coalesce(func.max(SmartDocumentVersion.version), 0)) .where( SmartDocumentVersion.document_id == document_id, SmartDocumentVersion.is_deleted == False, # noqa: E712 ) ) return (result.scalar() or 0) + 1 @classmethod async def create_version( cls, db: AsyncSession, document_id: str, content: Dict[str, Any], title: Optional[str] = None, change_summary: Optional[str] = None, user_id: Optional[str] = None, ) -> SmartDocumentVersion: import json version_num = await cls.get_next_version(db, document_id) content_size = len(json.dumps(content, ensure_ascii=False)) version = SmartDocumentVersion( document_id=document_id, version=version_num, content=content, title=title, change_summary=change_summary, content_size=content_size, sys_creator_id=user_id, ) db.add(version) await db.commit() await db.refresh(version) await cls._cleanup_old_versions(db, document_id) return version @classmethod async def _cleanup_old_versions(cls, db: AsyncSession, document_id: str) -> None: """保留最近 MAX_VERSIONS_PER_DOC 个版本,软删除更早的""" result = await db.execute( select(SmartDocumentVersion.id) .where( SmartDocumentVersion.document_id == document_id, SmartDocumentVersion.is_deleted == False, # noqa: E712 ) .order_by(SmartDocumentVersion.version.desc()) .offset(cls.MAX_VERSIONS_PER_DOC) ) old_ids = [row[0] for row in result.all()] if old_ids: await db.execute( sa_update(SmartDocumentVersion) .where(SmartDocumentVersion.id.in_(old_ids)) .values(is_deleted=True) ) await db.commit() @classmethod async def get_versions( cls, db: AsyncSession, document_id: str, page: int = 1, page_size: int = 20, ) -> Tuple[List[SmartDocumentVersion], int]: base = ( select(SmartDocumentVersion) .where( SmartDocumentVersion.document_id == document_id, SmartDocumentVersion.is_deleted == False, # noqa: E712 ) ) count_result = await db.execute( select(func.count()).select_from(base.subquery()) ) total = count_result.scalar() or 0 result = await db.execute( base.order_by(SmartDocumentVersion.version.desc()) .offset((page - 1) * page_size) .limit(page_size) ) return list(result.scalars().all()), total @classmethod async def get_version_detail( cls, db: AsyncSession, version_id: str ) -> Optional[SmartDocumentVersion]: result = await db.execute( select(SmartDocumentVersion) .where( SmartDocumentVersion.id == version_id, SmartDocumentVersion.is_deleted == False, # noqa: E712 ) ) return result.scalar_one_or_none() @classmethod async def get_version_by_number( cls, db: AsyncSession, document_id: str, version: int ) -> Optional[SmartDocumentVersion]: result = await db.execute( select(SmartDocumentVersion) .where( SmartDocumentVersion.document_id == document_id, SmartDocumentVersion.version == version, SmartDocumentVersion.is_deleted == False, # noqa: E712 ) ) return result.scalar_one_or_none() @classmethod async def restore_version( cls, db: AsyncSession, document_id: str, version_id: str, user_id: Optional[str] = None ) -> Optional[SmartDocumentVersion]: """恢复到指定版本:先快照当前内容,再将目标版本内容写回文档""" version = await cls.get_version_detail(db, version_id) if not version or version.document_id != document_id: return None table = await db.execute( select(SmartTable) .where(SmartTable.id == document_id, SmartTable.is_deleted == False) # noqa: E712 ) doc = table.scalar_one_or_none() if not doc: return None if doc.content: await cls.create_version( db, document_id, doc.content, title=doc.name, change_summary=f"恢复版本 {version.version} 前的自动快照", user_id=user_id, ) doc.content = version.content doc.sys_modifier_id = user_id await db.commit() new_version = await cls.create_version( db, document_id, version.content, title=doc.name, change_summary=f"恢复至版本 {version.version}", user_id=user_id, ) return new_version @classmethod async def delete_version(cls, db: AsyncSession, version_id: str) -> bool: result = await db.execute( select(SmartDocumentVersion) .where( SmartDocumentVersion.id == version_id, SmartDocumentVersion.is_deleted == False, # noqa: E712 ) ) version = result.scalar_one_or_none() if not version: return False version.is_deleted = True await db.commit() return True class SmartDocumentTemplateService: """文档模板服务""" @classmethod async def get_list( cls, db: AsyncSession, category: Optional[str] = None, keyword: Optional[str] = None, page: int = 1, page_size: int = 50, ) -> Tuple[List[SmartDocumentTemplate], int]: base = select(SmartDocumentTemplate).where( SmartDocumentTemplate.is_deleted == False # noqa: E712 ) if category: base = base.where(SmartDocumentTemplate.category == category) if keyword: base = base.where( or_( SmartDocumentTemplate.name.ilike(f"%{keyword}%"), SmartDocumentTemplate.description.ilike(f"%{keyword}%"), ) ) count_result = await db.execute( select(func.count()).select_from(base.subquery()) ) total = count_result.scalar() or 0 result = await db.execute( base.order_by( SmartDocumentTemplate.is_system.desc(), SmartDocumentTemplate.use_count.desc(), SmartDocumentTemplate.sys_create_datetime.desc(), ) .offset((page - 1) * page_size) .limit(page_size) ) return list(result.scalars().all()), total @classmethod async def get_by_id( cls, db: AsyncSession, template_id: str ) -> Optional[SmartDocumentTemplate]: result = await db.execute( select(SmartDocumentTemplate) .where( SmartDocumentTemplate.id == template_id, SmartDocumentTemplate.is_deleted == False, # noqa: E712 ) ) return result.scalar_one_or_none() @classmethod async def create( cls, db: AsyncSession, name: str, content: Dict[str, Any], description: Optional[str] = None, icon: str = "FileText", category: str = "custom", preview_image: Optional[str] = None, user_id: Optional[str] = None, ) -> SmartDocumentTemplate: template = SmartDocumentTemplate( name=name, description=description, icon=icon, category=category, content=content, preview_image=preview_image, is_system=False, sys_creator_id=user_id, ) db.add(template) await db.commit() await db.refresh(template) return template @classmethod async def update( cls, db: AsyncSession, template_id: str, data: Dict[str, Any], user_id: Optional[str] = None, ) -> Optional[SmartDocumentTemplate]: template = await cls.get_by_id(db, template_id) if not template: return None for key, value in data.items(): if value is not None and hasattr(template, key): setattr(template, key, value) template.sys_modifier_id = user_id await db.commit() await db.refresh(template) return template @classmethod async def delete(cls, db: AsyncSession, template_id: str) -> bool: template = await cls.get_by_id(db, template_id) if not template: return False if template.is_system: return False template.is_deleted = True await db.commit() return True @classmethod async def increment_use_count(cls, db: AsyncSession, template_id: str) -> None: await db.execute( sa_update(SmartDocumentTemplate) .where(SmartDocumentTemplate.id == template_id) .values(use_count=SmartDocumentTemplate.use_count + 1) ) await db.commit() @classmethod async def get_categories(cls, db: AsyncSession) -> List[str]: result = await db.execute( select(SmartDocumentTemplate.category) .where(SmartDocumentTemplate.is_deleted == False) # noqa: E712 .distinct() ) return [row[0] for row in result.all()] @classmethod async def create_from_document( cls, db: AsyncSession, document_id: str, name: str, description: Optional[str] = None, category: str = "custom", user_id: Optional[str] = None, ) -> Optional[SmartDocumentTemplate]: """从现有文档创建模板""" doc_result = await db.execute( select(SmartTable) .where(SmartTable.id == document_id, SmartTable.is_deleted == False) # noqa: E712 ) doc = doc_result.scalar_one_or_none() if not doc or getattr(doc, "type", "table") != "document" or not doc.content: return None return await cls.create( db, name=name, content=doc.content, description=description, category=category, user_id=user_id, ) class WikiSpaceService(BaseService[SmartWikiSpace, WikiSpaceCreate, WikiSpaceUpdate]): model = SmartWikiSpace RESOURCE_TYPE = "wiki_space" @classmethod async def get_list_with_count( cls, db: AsyncSession, page: int = 1, page_size: int = 100, user_id: Optional[str] = None, is_superuser: bool = False, ) -> Tuple[List[dict], int]: """获取文档库列表,附带每个库的文档数量。非超管只能看到自己创建的或 public/team 的""" base_query = select(SmartWikiSpace).where(SmartWikiSpace.is_deleted == False) # noqa: E712 if not is_superuser and user_id: base_query = base_query.where( or_( SmartWikiSpace.sys_creator_id == user_id, SmartWikiSpace.visibility.in_(["public", "team"]), ) ) elif not is_superuser: base_query = base_query.where( SmartWikiSpace.visibility.in_(["public", "team"]) ) count_result = await db.execute( select(func.count()).select_from(base_query.subquery()) ) total = count_result.scalar() or 0 offset = (page - 1) * page_size spaces_result = await db.execute( base_query.order_by(desc(SmartWikiSpace.sort), desc(SmartWikiSpace.sys_create_datetime)) .offset(offset) .limit(page_size) ) spaces = list(spaces_result.scalars().all()) result = [] for space in spaces: count_q = await db.execute( select(func.count(SmartTable.id)) .where( SmartTable.wiki_space_id == space.id, SmartTable.is_deleted == False, # noqa: E712 ) ) doc_count = count_q.scalar() or 0 result.append({ "space": space, "document_count": doc_count, }) return result, total @classmethod async def get_documents( cls, db: AsyncSession, space_id: str, user_id: Optional[str] = None, dept_id: Optional[str] = None, role_ids: Optional[List[str]] = None, is_superuser: bool = False, ) -> List[SmartTable]: """获取文档库内的文档(按 sort 排序),非超管只能看到自己有权限的""" from zq_smart_table.permission.model import SmartTableCollaborator base_query = ( select(SmartTable) .where( SmartTable.wiki_space_id == space_id, SmartTable.is_deleted == False, # noqa: E712 ) ) if not is_superuser and user_id: subject_conds = [ and_( SmartTableCollaborator.subject_type == "user", SmartTableCollaborator.subject_id == user_id, ) ] if dept_id: subject_conds.append( and_( SmartTableCollaborator.subject_type == "dept", SmartTableCollaborator.subject_id == dept_id, ) ) if role_ids: for rid in role_ids: subject_conds.append( and_( SmartTableCollaborator.subject_type == "role", SmartTableCollaborator.subject_id == rid, ) ) collab_table_ids = ( select(SmartTableCollaborator.table_id) .where( SmartTableCollaborator.is_deleted == False, # noqa: E712 or_(*subject_conds), ) .distinct() ) base_query = base_query.where( or_( SmartTable.sys_creator_id == user_id, SmartTable.id.in_(collab_table_ids), ) ) result = await db.execute( base_query.order_by(SmartTable.sort, SmartTable.sys_create_datetime) ) return list(result.scalars().all()) @classmethod async def check_space_access( cls, space: SmartWikiSpace, user_id: Optional[str] = None, is_superuser: bool = False, ) -> bool: """检查用户是否有权访问该文档库""" if is_superuser: return True if space.visibility in ("public", "team"): return True if user_id and space.sys_creator_id == user_id: return True return False @classmethod async def add_document( cls, db: AsyncSession, space_id: str, name: str, parent_id: Optional[str] = None, content: Optional[dict] = None, user_id: Optional[str] = None, ) -> SmartTable: """在文档库中创建文档""" max_sort_q = await db.execute( select(func.coalesce(func.max(SmartTable.sort), 0)) .where( SmartTable.wiki_space_id == space_id, SmartTable.is_deleted == False, # noqa: E712 ) ) next_sort = (max_sort_q.scalar() or 0) + 1 doc = SmartTable( name=name, icon="FileText", type="document", wiki_space_id=space_id, parent_id=parent_id, content=content, sort=next_sort, ) if user_id: doc.sys_creator_id = user_id doc.sys_modifier_id = user_id db.add(doc) await db.commit() await db.refresh(doc) return doc @classmethod async def delete_with_documents(cls, db: AsyncSession, space_id: str) -> bool: """软删除文档库及其所有文档""" space = await cls.get_by_id(db, space_id) if not space: return False await db.execute( sa_update(SmartTable) .where( SmartTable.wiki_space_id == space_id, SmartTable.is_deleted == False, # noqa: E712 ) .values(is_deleted=True) ) await cls.delete(db, space_id) return True