316 lines
10 KiB
Python
316 lines
10 KiB
Python
"""
|
||||
|
|
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)})"
|