已合并
fix:建表的时候,也要建立索引 #345
王欣创建于 6月26日
fix:建表的时候,也要建立索引 #345
已合并
共 2 个文件变更+51-9
| @@ -5,14 +5,14 @@ from datetime import datetime | |||
| 5 | import logging | 5 | import logging |
| 6 | from typing import Optional, Any | 6 | from typing import Optional, Any |
| 7 | import json | 7 | import json |
| 8 | -from sqlalchemy import Column, Integer, String, DateTime, JSON, Boolean, Float, create_engine, text, inspect | 8 | +from sqlalchemy import Column, Integer, String, DateTime, JSON, Boolean, Float, create_engine, text, inspect, Index |
| 9 | from sqlalchemy.ext.asyncio import AsyncSession, create_async_engine, async_sessionmaker | 9 | from sqlalchemy.ext.asyncio import AsyncSession, create_async_engine, async_sessionmaker |
| 10 | from sqlalchemy.orm import DeclarativeBase | 10 | from sqlalchemy.orm import DeclarativeBase |
| 11 | from sqlalchemy import select, update, delete, func | 11 | from sqlalchemy import select, update, delete, func |
| 12 | 12 | ||
| 13 | from ..log import get_logger | 13 | from ..log import get_logger |
| 14 | from .handler import DBHandler | 14 | from .handler import DBHandler |
| 15 | -from .table_def import TableDefinition, ColumnDefinition | 15 | +from .table_def import TableDefinition, ColumnDefinition, IndexDefinition |
| 16 | 16 | ||
| 17 | logger = get_logger(__name__) | 17 | logger = get_logger(__name__) |
| 18 | 18 | ||
| @@ -88,9 +88,17 @@ class SQLAlchemyHandler(DBHandler): | |||
| 88 | return String(length) | 88 | return String(length) |
| 89 | return sa_type | 89 | return sa_type |
| 90 | 90 | ||
| 91 | - @staticmethod | 91 | + 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] = table | 240 | 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 | ||