""" Schema数据模型定义 """ from dataclasses import dataclass, field from typing import List, Optional, Dict, Any from datetime import datetime @dataclass class Column: """列字段定义""" name: str data_type: str comment: Optional[str] = None nullable: bool = True is_primary_key: bool = False is_foreign_key: bool = False def __str__(self): null_str = "NULL" if self.nullable else "NOT NULL" pk_str = " PK" if self.is_primary_key else "" comment_str = f" -- {self.comment}" if self.comment else "" return f"{self.name} {self.data_type}{null_str}{pk_str}{comment_str}" @dataclass class ForeignKey: """外键关系""" columns: List[str] # 本表字段 ref_table: str # 引用表 ref_columns: List[str] # 引用字段 @dataclass class Table: """数据表定义""" name: str comment: Optional[str] = None columns: List[Column] = field(default_factory=list) primary_keys: List[str] = field(default_factory=list) foreign_keys: List[ForeignKey] = field(default_factory=list) # 向量embedding(可选,用于检索) embedding: Optional[Any] = None @property def column_dict(self) -> Dict[str, Column]: """字段名到Column对象的映射""" return {col.name: col for col in self.columns} @property def column_names(self) -> List[str]: """所有字段名列表""" return [col.name for col in self.columns] def to_compact_string(self, max_columns: int = 20) -> str: """ 生成紧凑的Schema字符串(用于LLM输入) Args: max_columns: 最大字段数,超过时截断 Returns: 紧凑的Schema描述字符串 """ parts = [f"TABLE {self.name} ("] # 只显示关键字段(主键、常见字段) display_columns = self.columns[:max_columns] if len(self.columns) > max_columns else self.columns col_strs = [] for col in display_columns: col_str = f" {col.name}: {col.data_type}" if col.is_primary_key: col_str += " [PK]" if col.comment: col_str += f" -- {col.comment}" col_strs.append(col_str) parts.append(",\n".join(col_strs)) if len(self.columns) > max_columns: parts.append(f"\n -- 还有 {len(self.columns) - max_columns} 个字段未显示") parts.append(")") return "".join(parts) def to_dict(self) -> Dict: """转换为字典格式""" return { "name": self.name, "comment": self.comment, "columns": [ { "name": col.name, "data_type": col.data_type, "comment": col.comment, "nullable": col.nullable, "is_primary_key": col.is_primary_key, } for col in self.columns ], "primary_keys": self.primary_keys, "foreign_keys": [ { "columns": fk.columns, "ref_table": fk.ref_table, "ref_columns": fk.ref_columns, } for fk in self.foreign_keys ], } @dataclass class DatabaseSchema: """数据库Schema(多个表的集合)""" name: str tables: List[Table] = field(default_factory=list) created_at: datetime = field(default_factory=datetime.now) @property def table_dict(self) -> Dict[str, Table]: """表名字典""" return {tbl.name: tbl for tbl in self.tables} def get_table(self, table_name: str) -> Optional[Table]: """获取指定表""" return self.table_dict.get(table_name) def to_summary_string(self, include_columns: bool = True) -> str: """ 生成Schema摘要字符串(用于LLM输入) Args: include_columns: 是否包含字段信息 Returns: Schema描述字符串 """ lines = [f"数据库: {self.name}\n"] lines.append(f"表总数: {len(self.tables)}\n") lines.append("=" * 60 + "\n") for table in self.tables: lines.append(f"表名: {table.name}") if table.comment: lines.append(f"描述: {table.comment}") lines.append(f"字段数: {len(table.columns)}") if include_columns and table.columns: lines.append("字段列表:") for col in table.columns[:30]: # 限制字段数 pk_marker = " [PK]" if col.is_primary_key else "" null_marker = " NULL" if col.nullable else " NOT NULL" 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) > 30: lines.append(f" ... 还有 {len(table.columns) - 30} 个字段") # 外键关系 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_compact_dict(self) -> Dict: """紧凑字典格式(用于序列化)""" return { "database": self.name, "tables": [tbl.to_dict() for tbl in self.tables], }