""" Schema管理器 - 管理数据库Schema的加载、索引和检索 """ import json import hashlib from pathlib import Path from typing import List, Dict, Optional, Tuple import logging from .models import DatabaseSchema, Table from .loader import SchemaLoader, load_schema_from_json, load_schema_from_g3sb logger = logging.getLogger(__name__) class SchemaManager: """ Schema管理器 功能: 1. 加载和解析Schema文件 2. 管理表关系 3. 提供表检索接口 4. 生成紧凑的Schema描述 """ def __init__(self, schema: DatabaseSchema): """ 初始化Schema管理器 Args: schema: DatabaseSchema对象 """ self.schema = schema self._table_dict = {tbl.name: tbl for tbl in schema.tables} self._embedding_index = None # 向量索引(延迟初始化) self._cache = {} @classmethod def load_from_json( cls, json_path: str, g3sb_meta_path: Optional[str] = None ) -> "SchemaManager": """ 从JSON文件加载Schema Args: json_path: JSON文件路径 g3sb_meta_path: 可选,G3SB table_meta.json(与 table_structure 配对) Returns: SchemaManager实例 """ schema = load_schema_from_json(json_path, g3sb_meta_path=g3sb_meta_path) return cls(schema) @classmethod def load_from_g3sb(cls, meta_path: str, structure_path: str) -> "SchemaManager": """ 从G3SB格式加载Schema Args: meta_path: table_meta.json路径 structure_path: table_structure.json路径 Returns: SchemaManager实例 """ schema = load_schema_from_g3sb(meta_path, structure_path) return cls(schema) @classmethod def from_dict(cls, data: Dict) -> "SchemaManager": """ 从字典创建SchemaManager Args: data: 包含database和tables的字典 Returns: SchemaManager实例 """ loader = SchemaLoader() schema = loader._parse_table(data) if isinstance(data, dict) else None if not schema: raise ValueError("Invalid schema data") return cls(schema) # ========== 查询接口 ========== def get_table(self, table_name: str) -> Optional[Table]: """获取指定表""" return self._table_dict.get(table_name) def list_tables(self) -> List[str]: """获取所有表名""" return list(self._table_dict.keys()) def get_tables(self) -> List[Table]: """获取所有Table对象""" return list(self._table_dict.values()) def get_related_tables(self, table_name: str, depth: int = 1) -> List[Table]: """ 获取关联表(通过外键关系) Args: table_name: 起始表名 depth: 递归深度(1=直接关联,2=间接关联...) Returns: 相关表列表 """ result = set() visited = set() def _traverse(tname: str, current_depth: int): if tname not in self._table_dict or current_depth > depth: return if tname in visited: return visited.add(tname) table = self._table_dict[tname] for fk in table.foreign_keys: result.add(fk.ref_table) _traverse(fk.ref_table, current_depth + 1) # 反向外键(被引用的表) for other_table in self._table_dict.values(): for fk in other_table.foreign_keys: if fk.ref_table == tname: result.add(other_table.name) _traverse(other_table.name, current_depth + 1) _traverse(table_name, 0) return [self._table_dict[t] for t in result if t in self._table_dict] # ========== Schema生成 ========== def to_compact_string( self, table_names: List[str] = None, include_columns: bool = True, max_columns_per_table: int = 15, ) -> str: """ 生成紧凑的Schema描述字符串(用于LLM输入) Args: table_names: 指定包含的表(None表示所有表) include_columns: 是否包含字段详情 max_columns_per_table: 每表最多显示字段数 Returns: Schema描述字符串 """ if table_names is None: tables = self.get_tables() else: missing = [n for n in table_names if n not in self._table_dict] if missing: logger.warning( "以下表名不在已加载Schema中,已从本次Schema片段中省略: %s", missing, ) tables = [t for t in self.get_tables() if t.name in table_names] lines = [] lines.append(f"数据库: {self.schema.name}") lines.append(f"涉及表数: {len(tables)}") lines.append("=" * 60) lines.append("") for table in tables: lines.append(f"【表】{table.name}") if table.comment: lines.append(f" 描述: {table.comment}") if include_columns and table.columns: lines.append(f" 字段:") for col in table.columns[:max_columns_per_table]: pk_marker = " [PK]" if col.is_primary_key else "" null_marker = " NOT NULL" if not col.nullable else "" comment = f" -- {col.comment}" if col.comment else "" lines.append(f" {col.name}: {col.data_type}{pk_marker}{null_marker}{comment}") if len(table.columns) > max_columns_per_table: lines.append(f" ... 还有 {len(table.columns) - max_columns_per_table} 个字段") # 外键关系 if table.foreign_keys: lines.append(" 外键:") for fk in table.foreign_keys[:5]: lines.append( f" {', '.join(fk.columns)} → {fk.ref_table}({', '.join(fk.ref_columns)})" ) lines.append("") return "\n".join(lines) def to_filtered_schema(self, table_names: List[str]) -> DatabaseSchema: """ 根据表名过滤,生成子集Schema Args: table_names: 要保留的表名列表 Returns: 新的DatabaseSchema对象 """ filtered_tables = [t for t in self.get_tables() if t.name in table_names] return DatabaseSchema( name=f"{self.schema.name}_filtered", tables=filtered_tables, ) # ========== 向量索引支持 ========== def set_embeddings(self, table_embeddings: Dict[str, List[float]]): """ 为表设置预计算的embedding向量 Args: table_embeddings: {table_name: embedding_vector} """ for table in self.get_tables(): if table.name in table_embeddings: table.embedding = table_embeddings[table.name] def get_table_for_embedding( self, max_columns: int = 48, max_col_comment_chars: int = 160, max_total_chars: int = 8000, ) -> List[Dict]: """ 获取用于embedding的表信息列表(单条文本同时承载 meta + structure 的可检索信息)。 - 表级:表名、表注释(来自 table_meta 合并后的 Table.comment) - 列级:字段名、类型、主键标记、列注释(来自 table_structure 解析后的 Column) Returns: [{"name": "...", "text": "..."}, ...] """ result = [] for table in self.get_tables(): text_parts = [f"表名: {table.name}"] if table.comment: text_parts.append(f"描述: {table.comment}") col_segments: List[str] = [] shown = table.columns[:max_columns] for col in shown: seg = f"{col.name} {col.data_type}" if col.is_primary_key: seg += " PK" if col.comment: c = col.comment.strip().replace("\n", " ") if len(c) > max_col_comment_chars: c = c[: max_col_comment_chars - 1] + "…" seg += f" — {c}" col_segments.append(seg) if col_segments: text_parts.append("字段: " + ";".join(col_segments)) if len(table.columns) > max_columns: rest = len(table.columns) - max_columns text_parts.append(f"另有{rest}个字段未列出(共{len(table.columns)}列)") text = "。".join(text_parts) if len(text) > max_total_chars: text = text[: max_total_chars - 1] + "…" result.append({ "name": table.name, "text": text, }) return result # ========== 缓存支持 ========== def get_cached_result(self, key: str) -> Optional[str]: """获取缓存结果""" return self._cache.get(key) def set_cached_result(self, key: str, value: str, ttl: int = 3600): """设置缓存结果""" self._cache[key] = value # TODO: 可扩展为Redis缓存 # ========== 统计信息 ========== def get_statistics(self) -> Dict: """获取Schema统计信息""" total_columns = sum(len(t.columns) for t in self.get_tables()) tables_with_pk = sum(1 for t in self.get_tables() if t.primary_keys) tables_with_fk = sum(1 for t in self.get_tables() if t.foreign_keys) return { "database": self.schema.name, "total_tables": len(self.get_tables()), "total_columns": total_columns, "avg_columns_per_table": total_columns / len(self.get_tables()) if self.get_tables() else 0, "tables_with_primary_key": tables_with_pk, "tables_with_foreign_key": tables_with_fk, "table_names": self.list_tables(), } def __len__(self) -> int: return len(self.get_tables()) def __repr__(self) -> str: return f"SchemaManager(database={self.schema.name}, tables={len(self)})"