first commit

This commit is contained in:
陈辅元
2026-04-10 16:52:07 +08:00
commit 84fe545640
87 changed files with 20842 additions and 0 deletions
+1
View File
@@ -0,0 +1 @@
# schema 包初始化
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
+229
View File
@@ -0,0 +1,229 @@
"""
Schema向量索引构建器 - 基于ChromaDB
"""
import chromadb
from chromadb.config import Settings as ChromaSettings
from typing import Any, List, Dict, Optional
import logging
from pathlib import Path
from schema.manager import SchemaManager
logger = logging.getLogger(__name__)
class SchemaIndexer:
"""
Schema向量索引器
功能:
1. 为表结构构建向量索引
2. 基于问题的表检索
3. 持久化存储和加载
"""
def __init__(
self,
embedder: Any,
persist_dir: str = "./data/embeddings",
collection_name: str = "schema_tables",
):
"""
初始化索引器
Args:
embedder: Embedding模型实例
persist_dir: 向量数据库持久化目录
collection_name: 集合名称
"""
self.embedder = embedder
self.persist_dir = Path(persist_dir)
self.persist_dir.mkdir(parents=True, exist_ok=True)
# 初始化ChromaDB客户端
self.client = chromadb.PersistentClient(
path=str(self.persist_dir),
settings=ChromaSettings(anonymized_telemetry=False),
)
# 获取或创建集合
self.collection = self.client.get_or_create_collection(
name=collection_name,
metadata={"hnsw:space": "cosine"}, # 使用余弦相似度
)
logger.info(f"[OK] 初始化SchemaIndexer: collection={collection_name}")
def build_index(
self,
schema_manager: SchemaManager,
batch_size: int = 32,
force_rebuild: bool = False,
) -> bool:
"""
为Schema构建向量索引
Args:
schema_manager: SchemaManager实例
batch_size: 批处理大小
force_rebuild: 是否强制重建(默认为False,增量更新)
Returns:
True=成功,False=已存在且未强制重建
"""
existing_ids = set(self.collection.get()["ids"]) if self.collection.count() > 0 else set()
if not force_rebuild and existing_ids:
logger.info(f"索引已存在({len(existing_ids)}条记录),跳过构建")
return False
if force_rebuild and existing_ids:
logger.info(f"强制重建索引,删除{len(existing_ids)}条旧记录")
self.collection.delete()
# 准备数据
table_texts = []
table_ids = []
table_metadatas = []
for table_info in schema_manager.get_table_for_embedding():
table_texts.append(table_info["text"])
table_ids.append(table_info["name"])
table_metadatas.append({
"table_name": table_info["name"],
"column_count": len(schema_manager.get_table(table_info["name"]).columns),
})
# 批量计算embedding
logger.info(f"计算{len(table_texts)}张表的embedding...")
embeddings = self.embedder.encode(table_texts, batch_size=batch_size)
# 存入ChromaDB
self.collection.add(
embeddings=embeddings.tolist(),
documents=table_texts,
metadatas=table_metadatas,
ids=table_ids,
)
logger.info(f"[OK] 索引构建完成,共{len(table_ids)}张表")
return True
def search(
self,
query: str,
top_k: int = 20,
score_threshold: float = 0.3,
) -> List[Dict]:
"""
检索相关表
Args:
query: 查询文本(用户问题)
top_k: 返回前K个结果
score_threshold: 相似度阈值(0-1),低于此值的结果会被过滤
Returns:
检索结果列表,按相似度降序排列
[
{
"table_name": "表名",
"score": 0.95,
"document": "表描述文本",
"metadata": {...}
},
...
]
"""
# 编码查询文本
query_embedding = self.embedder.encode([query])
# 执行检索
results = self.collection.query(
query_embeddings=query_embedding.tolist(),
n_results=min(top_k, self.collection.count()),
)
# 格式化结果
formatted = []
if results["ids"] and len(results["ids"][0]) > 0:
for idx, (table_id, distance, metadata, document) in enumerate(zip(
results["ids"][0],
results["distances"][0],
results["metadatas"][0],
results["documents"][0],
)):
score = 1 - distance # 余弦相似度(已归一化,转换为0-1,越大越相似)
if score >= score_threshold:
formatted.append({
"table_name": table_id,
"score": float(score),
"document": document,
"metadata": metadata,
"rank": idx + 1,
})
logger.debug(f"检索 '{query[:50]}...' -> 找到{len(formatted)}个相关表(阈值={score_threshold})")
return formatted
def search_by_table_names(self, table_names: List[str]) -> List[Dict]:
"""
直接通过表名获取表信息
Args:
table_names: 表名列表
Returns:
表信息列表
"""
results = []
for name in table_names:
try:
result = self.collection.get(ids=[name], include=["metadatas", "documents"])
if result["ids"]:
results.append({
"table_name": name,
"score": 1.0,
"document": result["documents"][0],
"metadata": result["metadatas"][0],
})
except Exception as e:
logger.warning(f"获取表信息失败 {name}: {e}")
return results
def get_all_tables(self) -> List[str]:
"""获取索引中的所有表名"""
result = self.collection.get()
return result["ids"] if result["ids"] else []
def delete_table(self, table_name: str):
"""从索引中删除表"""
self.collection.delete(ids=[table_name])
logger.info(f"已删除表索引: {table_name}")
def clear(self):
"""清空索引"""
self.collection.delete()
logger.info("索引已清空")
def count(self) -> int:
"""获取索引中的表数量"""
return self.collection.count()
# ========== 统计与调试 ==========
def get_statistics(self) -> Dict:
"""获取索引统计信息"""
count = self.collection.count()
result = self.collection.get(include=["metadatas"])
total_columns = sum(m.get("column_count", 0) for m in result["metadatas"])
return {
"indexed_tables": count,
"total_columns": total_columns,
"avg_columns": total_columns / count if count > 0 else 0,
"persist_dir": str(self.persist_dir),
}
+319
View File
@@ -0,0 +1,319 @@
"""
Schema加载器 - 解析JSON/DDL格式的数据库Schema
"""
import json
import re
from pathlib import Path
from typing import List, Dict, Optional, Tuple
import logging
from .models import Table, Column, ForeignKey, DatabaseSchema
logger = logging.getLogger(__name__)
class SchemaLoader:
"""Schema加载器 - 支持JSON和DDL格式"""
def __init__(self, schema_dir: str = None):
"""
初始化Schema加载器
Args:
schema_dir: Schema文件目录路径
"""
self.schema_dir = Path(schema_dir) if schema_dir else None
def load_from_json(
self, json_path: str, g3sb_meta_path: Optional[str] = None
) -> DatabaseSchema:
"""
从JSON文件加载Schema
JSON格式示例:
{
"database": "db_name",
"tables": [
{
"name": "table1",
"comment": "表描述",
"columns": [
{"name": "col1", "type": "INT", "comment": "...", "nullable": false}
],
"primary_keys": ["col1"],
"foreign_keys": [
{"columns": ["col2"], "ref_table": "table2", "ref_columns": ["col1"]}
]
}
]
}
亦支持 G3SB table_structure.json:顶层为 ``schemas`` 字典(与 ``tables`` 数组二选一
时优先使用非空的 ``schemas``)。可选 ``g3sb_meta_path`` 指向 table_meta.json
以合并表级注释。
Args:
json_path: JSON文件路径
g3sb_meta_path: 可选,G3SB table_meta.json(含 tables.{表名}.comment)
Returns:
DatabaseSchema对象
"""
path = Path(json_path)
if not path.exists():
raise FileNotFoundError(f"Schema文件不存在: {json_path}")
with open(path, 'r', encoding='utf-8') as f:
data = json.load(f)
# 提取数据库名
db_name = data.get("database", path.stem)
schemas_map = data.get("schemas")
tables_data = data.get("tables")
# 优先非空 schemas(G3SB 结构字典),避免误把其它 truthy 的 tables 键当好列表解析
if isinstance(schemas_map, dict) and schemas_map:
table_meta: Dict = {}
if g3sb_meta_path:
mp = Path(g3sb_meta_path)
if mp.is_file():
with open(mp, 'r', encoding='utf-8') as mf:
meta_data = json.load(mf)
table_meta = meta_data.get("tables", {}) or {}
else:
logger.warning(
"G3SB meta 文件不存在,将跳过表注释: %s", g3sb_meta_path
)
tables = []
for table_name, structure_str in schemas_map.items():
if not isinstance(structure_str, str):
continue
comment = None
if table_meta and isinstance(table_meta.get(table_name), dict):
comment = table_meta[table_name].get("comment")
columns = self._parse_g3sb_table_structure(structure_str)
tables.append(
Table(
name=table_name,
comment=comment,
columns=columns,
primary_keys=self._extract_primary_keys(columns),
foreign_keys=[],
)
)
logger.info(
"[OK] 加载Schema完成(G3SB schemas): %s, 共%d张表", db_name, len(tables)
)
return DatabaseSchema(name=db_name, tables=tables)
if isinstance(tables_data, list) and tables_data:
tables = [self._parse_table(table_data) for table_data in tables_data]
logger.info(f"[OK] 加载Schema完成: {db_name}, 共{len(tables)}张表")
return DatabaseSchema(name=db_name, tables=tables)
tables = []
logger.info(f"[OK] 加载Schema完成: {db_name}, 共{len(tables)}张表")
return DatabaseSchema(name=db_name, tables=tables)
def load_from_g3sb_format(self, meta_path: str, structure_path: str) -> DatabaseSchema:
"""
加载G3SB系统的数据字典格式(两个JSON文件)
Args:
meta_path: table_meta.json路径(表注释)
structure_path: table_structure.json路径(表结构)
Returns:
DatabaseSchema对象
"""
# 1. 加载表元数据(表注释)
with open(meta_path, 'r', encoding='utf-8') as f:
meta_data = json.load(f)
table_meta = meta_data.get("tables", {})
# 2. 加载表结构
with open(structure_path, 'r', encoding='utf-8') as f:
structure_data = json.load(f)
table_structures = structure_data.get("schemas", {})
# 3. 解析所有表
tables = []
for table_name, structure_str in table_structures.items():
comment = table_meta.get(table_name, {}).get("comment")
# 解析表结构字符串
# 格式: "TABLE TableName (col1:TYPE -- comment, col2:TYPE, ...)"
columns = self._parse_g3sb_table_structure(structure_str)
table = Table(
name=table_name,
comment=comment,
columns=columns,
primary_keys=self._extract_primary_keys(columns),
foreign_keys=[] # G3SB格式没有外键信息,后续可补充
)
tables.append(table)
logger.info(f"[OK] 加载G3SB Schema完成: 共{len(tables)}张表")
return DatabaseSchema(name="G3SB_DB", tables=tables)
def _parse_table(self, data: Dict) -> Table:
"""解析单表JSON数据"""
columns = []
for col_data in data.get("columns", []):
col = Column(
name=col_data["name"],
data_type=col_data["type"],
comment=col_data.get("comment"),
nullable=col_data.get("nullable", True),
is_primary_key=col_data.get("is_primary_key", False),
)
columns.append(col)
# 外键解析
foreign_keys = []
for fk_data in data.get("foreign_keys", []):
fk = ForeignKey(
columns=fk_data["columns"],
ref_table=fk_data["ref_table"],
ref_columns=fk_data["ref_columns"],
)
foreign_keys.append(fk)
return Table(
name=data["name"],
comment=data.get("comment"),
columns=columns,
primary_keys=data.get("primary_keys", []),
foreign_keys=foreign_keys,
)
def _parse_g3sb_table_structure(self, structure_str: str) -> List[Column]:
"""
解析G3SB格式的表结构字符串
示例:
"TABLE BCAccountCash (AccountID:NCHAR, RegionID:NCHAR, CurrencyID:NCHAR, Settled:DECIMAL -- Settled balance, ...)"
"""
# 提取括号内的内容
match = re.search(r'\((.*)\)', structure_str)
if not match:
logger.warning(f"无法解析表结构: {structure_str[:100]}")
return []
inner = match.group(1)
columns = []
# 按逗号分割字段(注意注释中可能包含逗号)
parts = self._split_columns(inner)
for part in parts:
part = part.strip()
if not part:
continue
# 解析字段定义:name:TYPE [-- comment]
# 支持格式: "FieldName:DATATYPE" 或 "FieldName:DATATYPE -- comment"
col_match = re.match(r'^(\w+)\s*:\s*([A-Za-z0-9()]+)', part)
if not col_match:
continue
col_name = col_match.group(1).strip()
col_type = col_match.group(2).strip()
# 提取注释(ASCII -- 或 G3SB 常用的 Unicode 长破折号 — U+2014)
comment = None
comment_match = re.search(r'(?:--|\u2014)\s*(.+)', part)
if comment_match:
comment = comment_match.group(1).strip()
# 判断是否可为空(通常有默认值或未标注NOT NULL即为NULL)
nullable = True # G3SB格式默认允许NULL
column = Column(
name=col_name,
data_type=col_type,
comment=comment,
nullable=nullable,
)
columns.append(column)
return columns
def _split_columns(self, inner: str) -> List[str]:
"""
按字段边界拆分(G3SB 注释里常有英文逗号,且用 — 而非 --)。
仅在「后面紧跟 标识符: 」的逗号处切分,这样注释内的逗号不会误拆列。
"""
if not inner or not inner.strip():
return []
# 下一列以 Name:TYPE 开头;避免在括号嵌套里误匹配可再收紧(当前 G3SB 类型无顶层逗号)
parts = re.split(r",\s*(?=\w+\s*:)", inner)
return [p.strip() for p in parts if p.strip()]
def _extract_primary_keys(self, columns: List[Column]) -> List[str]:
"""从字段列表中提取主键(简单启发式:字段名包含ID或明确标记)"""
pk_candidates = []
for col in columns:
# 简单规则:字段名以ID结尾,或名称包含key/id
if col.name.upper().endswith('ID') or 'KEY' in col.name.upper():
pk_candidates.append(col.name)
return pk_candidates[:1] # 暂时只取一个主键(简化)
def load_all_schemas(self) -> List[DatabaseSchema]:
"""
加载schema_dir下的所有Schema文件
Returns:
DatabaseSchema列表
"""
if not self.schema_dir:
raise ValueError("未指定schema_dir")
schemas = []
for json_file in self.schema_dir.glob("*.json"):
try:
schema = self.load_from_json(str(json_file))
schemas.append(schema)
except Exception as e:
logger.error(f"加载Schema失败 {json_file}: {e}")
logger.info(f"[OK] 共加载{len(schemas)}个Schema")
return schemas
# 便捷函数
def load_schema_from_g3sb(meta_path: str, structure_path: str) -> DatabaseSchema:
"""
从G3SB格式加载Schema的便捷函数
Args:
meta_path: table_meta.json路径
structure_path: table_structure.json路径
Returns:
DatabaseSchema对象
"""
loader = SchemaLoader()
return loader.load_from_g3sb_format(meta_path, structure_path)
def load_schema_from_json(
json_path: str, g3sb_meta_path: Optional[str] = None
) -> DatabaseSchema:
"""
从标准JSON加载Schema的便捷函数
Args:
json_path: JSON文件路径
g3sb_meta_path: 可选,G3SB table_meta.json
Returns:
DatabaseSchema对象
"""
loader = SchemaLoader()
return loader.load_from_json(json_path, g3sb_meta_path=g3sb_meta_path)
+315
View File
@@ -0,0 +1,315 @@
"""
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)})"
+179
View File
@@ -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],
}