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