first commit
This commit is contained in:
@@ -0,0 +1 @@
|
||||
# schema 包初始化
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -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),
|
||||
}
|
||||
@@ -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)
|
||||
@@ -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)})"
|
||||
@@ -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