99from sqlalchemy .dialects import postgresql
1010from sqlalchemy .sql .elements import ClauseElement
1111
12+ from .indexing import DEFAULT_INDEX_ACCESS_METHOD , validate_index_access_method
13+
1214
1315def _quote_ident (name : str ) -> str :
1416 return '"' + name .replace ('"' , '""' ) + '"'
@@ -35,13 +37,15 @@ def __init__(
3537 * ,
3638 table_schema : str | None = None ,
3739 where : str | None = None ,
40+ am : str = DEFAULT_INDEX_ACCESS_METHOD ,
3841 ) -> None :
3942 self .index_name = index_name
4043 self .table_name = table_name
4144 self .expressions = expressions
4245 self .key_field = key_field
4346 self .table_schema = table_schema
4447 self .where = where
48+ self .am = validate_index_access_method (am )
4549
4650 @classmethod
4751 def create_bm25_index (
@@ -54,6 +58,7 @@ def create_bm25_index(
5458 key_field : str ,
5559 table_schema : str | None = None ,
5660 where : str | None = None ,
61+ am : str = DEFAULT_INDEX_ACCESS_METHOD ,
5762 ) -> MigrateOperation :
5863 return operations .invoke (
5964 cls (
@@ -63,11 +68,12 @@ def create_bm25_index(
6368 key_field ,
6469 table_schema = table_schema ,
6570 where = where ,
71+ am = am ,
6672 )
6773 )
6874
6975 def reverse (self ) -> MigrateOperation :
70- return DropBM25IndexOp (index_name = self .index_name , if_exists = True , schema = self .table_schema )
76+ return DropBM25IndexOp (index_name = self .index_name , if_exists = True , schema = self .table_schema , am = self . am )
7177
7278
7379@Operations .implementation_for (CreateBM25IndexOp )
@@ -76,7 +82,7 @@ def _create_bm25_index_impl(operations: Operations, operation: CreateBM25IndexOp
7682 sql = (
7783 f"CREATE INDEX { _quote_ident (operation .index_name )} "
7884 f"ON { _quote_qualified (operation .table_schema , operation .table_name )} "
79- f"USING bm25 ({ expressions_sql } ) WITH (key_field={ _quote_literal (operation .key_field )} )"
85+ f"USING { operation . am } ({ expressions_sql } ) WITH (key_field={ _quote_literal (operation .key_field )} )"
8086 )
8187 if operation .where is not None :
8288 sql += f" WHERE { operation .where } "
@@ -95,6 +101,8 @@ def _render_create_bm25_index_op(autogen_context, op: CreateBM25IndexOp) -> str:
95101 parts .append (f"table_schema={ op .table_schema !r} " )
96102 if op .where is not None :
97103 parts .append (f"where={ op .where !r} " )
104+ if op .am != DEFAULT_INDEX_ACCESS_METHOD :
105+ parts .append (f"am={ op .am !r} " )
98106 return f"op.create_bm25_index({ ', ' .join (parts )} )"
99107
100108
@@ -110,6 +118,7 @@ def __init__(
110118 expressions : list [str ] | None = None ,
111119 key_field : str | None = None ,
112120 where : str | None = None ,
121+ am : str = DEFAULT_INDEX_ACCESS_METHOD ,
113122 ) -> None :
114123 self .index_name = index_name
115124 self .if_exists = if_exists
@@ -118,6 +127,7 @@ def __init__(
118127 self .expressions = expressions
119128 self .key_field = key_field
120129 self .where = where
130+ self .am = validate_index_access_method (am )
121131
122132 @classmethod
123133 def drop_bm25_index (
@@ -131,6 +141,7 @@ def drop_bm25_index(
131141 expressions : list [str ] | None = None ,
132142 key_field : str | None = None ,
133143 where : str | None = None ,
144+ am : str = DEFAULT_INDEX_ACCESS_METHOD ,
134145 ) -> MigrateOperation :
135146 return operations .invoke (
136147 cls (
@@ -141,6 +152,7 @@ def drop_bm25_index(
141152 expressions = expressions ,
142153 key_field = key_field ,
143154 where = where ,
155+ am = am ,
144156 )
145157 )
146158
@@ -155,6 +167,7 @@ def reverse(self) -> MigrateOperation:
155167 key_field = self .key_field ,
156168 table_schema = self .schema ,
157169 where = self .where ,
170+ am = self .am ,
158171 )
159172
160173
@@ -177,6 +190,8 @@ def _render_drop_bm25_index_op(autogen_context, op: DropBM25IndexOp) -> str:
177190 parts .append (f"key_field={ op .key_field !r} " )
178191 if op .where is not None :
179192 parts .append (f"where={ op .where !r} " )
193+ if op .am != DEFAULT_INDEX_ACCESS_METHOD :
194+ parts .append (f"am={ op .am !r} " )
180195 return f"op.drop_bm25_index({ ', ' .join (parts )} )"
181196
182197
@@ -229,7 +244,7 @@ def _autogen_bm25_meta_indexes(
229244
230245
231246def _autogen_bm25_db_indexes (conn , effective_schemas : set [str ]) -> dict [tuple [str , str ], dict ]:
232- """Return {(schema, index_name): {table_name, expressions, key_field, where}} from pg_indexes."""
247+ """Return {(schema, index_name): {table_name, expressions, key_field, where, am }} from pg_indexes."""
233248 from .indexing import (
234249 _extract_key_field ,
235250 _extract_where_clause ,
@@ -249,6 +264,7 @@ def _autogen_bm25_db_indexes(conn, effective_schemas: set[str]) -> dict[tuple[st
249264 "expressions" : [],
250265 "key_field" : _normalize_reloption_value (row ["key_field" ]) or "" ,
251266 "where" : _extract_where_clause (str (row ["indexdef" ])),
267+ "am" : str (row ["amname" ]),
252268 },
253269 )
254270 index_entry ["expressions" ].append (str (row ["keydef" ]))
@@ -447,6 +463,7 @@ def _compare_bm25_indexes(autogen_context, upgrade_ops, schemas) -> PriorityDisp
447463 expressions = db ["expressions" ],
448464 key_field = db ["key_field" ],
449465 where = db .get ("where" ),
466+ am = db ["am" ],
450467 )
451468 )
452469
@@ -467,11 +484,15 @@ def _compare_bm25_indexes(autogen_context, upgrade_ops, schemas) -> PriorityDisp
467484 key_field = key_field ,
468485 table_schema = key [0 ],
469486 where = meta_where ,
487+ am = validate_index_access_method (index .dialect_options ["postgresql" ].get ("using" )),
470488 )
471489
472490 if key not in db_bm25 :
473491 upgrade_ops .ops .append (create_op )
474492 else :
493+ # The AM name is intentionally not compared: paradedb and bm25 are
494+ # aliases for the same access method, so an AM-only difference must
495+ # not produce migration churn.
475496 db = db_bm25 [key ]
476497 expressions_changed = _normalized_expression_list (db ["expressions" ]) != _normalized_expression_list (
477498 expressions
@@ -488,6 +509,7 @@ def _compare_bm25_indexes(autogen_context, upgrade_ops, schemas) -> PriorityDisp
488509 expressions = db ["expressions" ],
489510 key_field = db ["key_field" ],
490511 where = db .get ("where" ),
512+ am = db ["am" ],
491513 )
492514 )
493515 upgrade_ops .ops .append (create_op )
0 commit comments