Files
2026-04-14 10:28:22 +08:00

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],
}