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)})"
|