first commit
This commit is contained in:
@@ -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)})"
|
||||
Reference in New Issue
Block a user