0.1.1 暂存
This commit is contained in:
@@ -0,0 +1,179 @@
|
||||
"""
|
||||
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],
|
||||
}
|
||||
Reference in New Issue
Block a user