180 lines
5.6 KiB
Python
180 lines
5.6 KiB
Python
"""
|
|||
|
|
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],
|
||
|
|
}
|