Files
2026-04-14 10:28:22 +08:00

316 lines
10 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
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)})"