已合并
fix:建表的时候,也要建立索引 #345
fix:建表的时候,也要建立索引 #345
已合并
王欣创建于 6月26日
2 个文件变更+51-9
@@ -5,14 +5,14 @@ from datetime import datetime
5import logging5import logging
6from typing import Optional, Any6from typing import Optional, Any
7import json7import json
8-from sqlalchemy import Column, Integer, String, DateTime, JSON, Boolean, Float, create_engine, text, inspect8+from sqlalchemy import Column, Integer, String, DateTime, JSON, Boolean, Float, create_engine, text, inspect, Index
9from sqlalchemy.ext.asyncio import AsyncSession, create_async_engine, async_sessionmaker9from sqlalchemy.ext.asyncio import AsyncSession, create_async_engine, async_sessionmaker
10from sqlalchemy.orm import DeclarativeBase10from sqlalchemy.orm import DeclarativeBase
11from sqlalchemy import select, update, delete, func11from sqlalchemy import select, update, delete, func
12 12 
13from ..log import get_logger13from ..log import get_logger
14from .handler import DBHandler14from .handler import DBHandler
15-from .table_def import TableDefinition, ColumnDefinition15+from .table_def import TableDefinition, ColumnDefinition, IndexDefinition
16 16 
17logger = get_logger(__name__)17logger = get_logger(__name__)
18 18 
@@ -88,9 +88,17 @@ class SQLAlchemyHandler(DBHandler):
88 return String(length)88 return String(length)
89 return sa_type89 return sa_type
90 90 
91- @staticmethod91+ def _get_dialect_name(self) -> str:
92- def _quote_identifier(identifier: str) -> str:92+ if self.engine is not None:
93- return f'"{identifier}"'93+ return self.engine.dialect.name
94+ from sqlalchemy.engine import make_url
95+ return make_url(self.database_url).get_backend_name()
96+ 
97+ def _quote_identifier(self, identifier: str) -> str:
98+ if self._get_dialect_name() == "mysql":
99+ return "`" + identifier.replace("`", "``") + "`"
100+ escaped = identifier.replace('"', '""')
101+ return f'"{escaped}"'
94 102 
95 def _get_column_sql_type(self, col_def: ColumnDefinition) -> str:103 def _get_column_sql_type(self, col_def: ColumnDefinition) -> str:
96 data_type = col_def.data_type.lower()104 data_type = col_def.data_type.lower()
@@ -167,6 +175,35 @@ class SQLAlchemyHandler(DBHandler):
167 col_def.name,175 col_def.name,
168 )176 )
169 177 
178+ def _build_index_name(self, table_name: str, idx_def: IndexDefinition) -> str:
179+ if idx_def.name:
180+ return idx_def.name
181+ return f"ix_{table_name}_{'_'.join(idx_def.columns)}"
182+ 
183+ def _create_table_indexes(self, sync_conn, table_def: TableDefinition) -> None:
184+ """通过 SQLAlchemy Index 创建索引,由方言层生成各数据库兼容的 DDL。"""
185+ inspector = inspect(sync_conn)
186+ existing_indexes = {
187+ idx["name"]
188+ for idx in inspector.get_indexes(table_def.table_name)
189+ }
190+ table = self._table_models[table_def.table_name].__table__
191+ for idx_def in table_def.indexes:
192+ idx_name = self._build_index_name(table_def.table_name, idx_def)
193+ if idx_name in existing_indexes:
194+ continue
195+ index = Index(
196+ idx_name,
197+ *[table.c[col] for col in idx_def.columns],
198+ unique=idx_def.unique,
199+ )
200+ index.create(sync_conn)
201+ logger.debug(
202+ "Created index during table init: table=%s, index=%s",
203+ table_def.table_name,
204+ idx_name,
205+ )
206+ 
170 async def init_table(self, table_def: TableDefinition) -> None:207 async def init_table(self, table_def: TableDefinition) -> None:
171 """初始化表(存在则跳过,不存在则创建)"""208 """初始化表(存在则跳过,不存在则创建)"""
172 logger.debug("Initializing table: table_name=%s", table_def.table_name)209 logger.debug("Initializing table: table_name=%s", table_def.table_name)
@@ -203,11 +240,16 @@ class SQLAlchemyHandler(DBHandler):
203 self._table_models[table_def.table_name] = table240 self._table_models[table_def.table_name] = table
204 241 
205 async with self.engine.begin() as conn:242 async with self.engine.begin() as conn:
206- await conn.run_sync(243+ def init_sync(sync_conn):
207- lambda sync_conn: Base.metadata.create_all(244+ inspector = inspect(sync_conn)
245+ table_exists = table_def.table_name in inspector.get_table_names()
246+ Base.metadata.create_all(
208 sync_conn, tables=[table.__table__]247 sync_conn, tables=[table.__table__]
209 )248 )
210- )249+ if not table_exists:
250+ self._create_table_indexes(sync_conn, table_def)
251+ 
252+ await conn.run_sync(init_sync)
211 logger.debug("Table initialized: table_name=%s", table_def.table_name)253 logger.debug("Table initialized: table_name=%s", table_def.table_name)
212 254 
213 async def _get_session(self) -> AsyncSession:255 async def _get_session(self) -> AsyncSession:
@@ -31,7 +31,7 @@ class TestSQLiteHandler(unittest.IsolatedAsyncioTestCase):
31 ColumnDefinition("value", "string", length=255, nullable=True),31 ColumnDefinition("value", "string", length=255, nullable=True),
32 ],32 ],
33 indexes=[33 indexes=[
34- IndexDefinition(["name"], unique=True),34+ IndexDefinition(["name"], unique=False),
35 ],35 ],
36 )36 )
37 37