11from __future__ import annotations
22
3+ import re
34from dataclasses import dataclass
45from typing import Any
56
1011from sqlalchemy .sql .elements import ClauseElement , ColumnElement
1112from sqlalchemy .sql .visitors import InternalTraversal
1213
14+ from .errors import (
15+ DuplicateTokenizerAliasError ,
16+ InvalidArgumentError ,
17+ InvalidBM25FieldError ,
18+ InvalidKeyFieldError ,
19+ MissingKeyFieldError ,
20+ )
21+
1322
1423@dataclass (frozen = True )
1524class TokenizerSpec :
@@ -23,7 +32,7 @@ def render(self) -> str:
2332 return self .raw_sql
2433
2534 if self .name is None :
26- raise ValueError ("tokenizer name is required unless raw_sql is provided" )
35+ raise InvalidArgumentError ("tokenizer name is required unless raw_sql is provided" )
2736
2837 if not self .options :
2938 return f"pdb.{ self .name } ()"
@@ -135,28 +144,28 @@ def validate_bm25_index(index: Index) -> None:
135144 return
136145
137146 if not index .expressions :
138- raise ValueError ("BM25 indexes must include at least one BM25Field" )
147+ raise InvalidBM25FieldError ("BM25 indexes must include at least one BM25Field" )
139148
140149 if not all (isinstance (expr , BM25Field ) for expr in index .expressions ):
141- raise ValueError ("BM25 indexes must use BM25Field for every indexed field" )
150+ raise InvalidBM25FieldError ("BM25 indexes must use BM25Field for every indexed field" )
142151
143152 aliases : set [str ] = set ()
144153 for expr in index .expressions :
145154 tokenizer = expr .tokenizer
146155 if tokenizer is None or tokenizer .alias is None :
147156 continue
148157 if tokenizer .alias in aliases :
149- raise ValueError (f"Duplicate tokenizer alias '{ tokenizer .alias } ' in BM25 index" )
158+ raise DuplicateTokenizerAliasError (f"Duplicate tokenizer alias '{ tokenizer .alias } ' in BM25 index" )
150159 aliases .add (tokenizer .alias )
151160
152161 with_options = index .dialect_options ["postgresql" ].get ("with" ) or {}
153162 key_field = with_options .get ("key_field" )
154163 if not key_field :
155- raise ValueError ("BM25 indexes require postgresql_with={'key_field': '<column>'}" )
164+ raise MissingKeyFieldError ("BM25 indexes require postgresql_with={'key_field': '<column>'}" )
156165
157166 field_names = {_bm25_field_name (expr ) for expr in index .expressions }
158167 if key_field not in field_names :
159- raise ValueError (f"BM25 key_field '{ key_field } ' must match one of the indexed BM25Field columns" )
168+ raise InvalidKeyFieldError (f"BM25 key_field '{ key_field } ' must match one of the indexed BM25Field columns" )
160169
161170
162171@event .listens_for (Index , "before_create" )
@@ -172,6 +181,99 @@ class IndexMeta:
172181 aliases : dict [str , str ]
173182
174183
184+ _KEY_FIELD_RE = re .compile (r"key_field\s*=\s*'?\"?([^'\",)\s]+)\"?'?" , re .IGNORECASE )
185+ _ALIAS_RE = re .compile (r"alias\s*=\s*([A-Za-z_][A-Za-z0-9_]*)" , re .IGNORECASE )
186+ _CAST_FIELD_RE = re .compile (r"^\(*\"?([A-Za-z_][A-Za-z0-9_]*)\"?\)*\s*::\s*pdb\." , re .IGNORECASE )
187+ _PLAIN_FIELD_RE = re .compile (r'^\(*"?([A-Za-z_][A-Za-z0-9_]*)"?\)*$' )
188+
189+
190+ def _split_top_level_csv (expr : str ) -> list [str ]:
191+ parts : list [str ] = []
192+ current : list [str ] = []
193+ depth = 0
194+ in_single = False
195+ in_double = False
196+
197+ for ch in expr :
198+ if ch == "'" and not in_double :
199+ in_single = not in_single
200+ current .append (ch )
201+ continue
202+ if ch == '"' and not in_single :
203+ in_double = not in_double
204+ current .append (ch )
205+ continue
206+ if not in_single and not in_double :
207+ if ch == "(" :
208+ depth += 1
209+ elif ch == ")" :
210+ depth = max (0 , depth - 1 )
211+ elif ch == "," and depth == 0 :
212+ piece = "" .join (current ).strip ()
213+ if piece :
214+ parts .append (piece )
215+ current = []
216+ continue
217+ current .append (ch )
218+
219+ tail = "" .join (current ).strip ()
220+ if tail :
221+ parts .append (tail )
222+ return parts
223+
224+
225+ def _extract_bm25_field_list (indexdef : str ) -> list [str ]:
226+ marker = re .search (r"USING\s+bm25\s*\(" , indexdef , re .IGNORECASE )
227+ if marker is None :
228+ return []
229+
230+ start = marker .end ()
231+ depth = 1
232+ in_single = False
233+ in_double = False
234+ i = start
235+ while i < len (indexdef ):
236+ ch = indexdef [i ]
237+ if ch == "'" and not in_double :
238+ in_single = not in_single
239+ elif ch == '"' and not in_single :
240+ in_double = not in_double
241+ elif not in_single and not in_double :
242+ if ch == "(" :
243+ depth += 1
244+ elif ch == ")" :
245+ depth -= 1
246+ if depth == 0 :
247+ return _split_top_level_csv (indexdef [start :i ])
248+ i += 1
249+ return []
250+
251+
252+ def _extract_field_name (field_expr : str ) -> str | None :
253+ expr = field_expr .strip ()
254+ cast_match = _CAST_FIELD_RE .match (expr )
255+ if cast_match :
256+ return cast_match .group (1 )
257+ plain_match = _PLAIN_FIELD_RE .match (expr )
258+ if plain_match :
259+ return plain_match .group (1 )
260+ return None
261+
262+
263+ def _extract_key_field (indexdef : str ) -> str | None :
264+ match = _KEY_FIELD_RE .search (indexdef )
265+ if match :
266+ return match .group (1 )
267+ return None
268+
269+
270+ def _extract_alias (index_expr : str ) -> str | None :
271+ match = _ALIAS_RE .search (index_expr )
272+ if match :
273+ return match .group (1 )
274+ return None
275+
276+
175277def describe (engine : Engine , table ) -> list [IndexMeta ]:
176278 query = text (
177279 """
@@ -188,21 +290,26 @@ def describe(engine: Engine, table) -> list[IndexMeta]:
188290 output : list [IndexMeta ] = []
189291 for row in rows :
190292 indexdef : str = row .indexdef
191- key_field : str | None = None
192- marker = "key_field='"
193- marker_idx = indexdef .find (marker )
194- if marker_idx != - 1 :
195- key_start = marker_idx + len (marker )
196- key_end = indexdef .find ("'" , key_start )
197- if key_end != - 1 :
198- key_field = indexdef [key_start :key_end ]
293+ key_field = _extract_key_field (indexdef )
294+ raw_fields = _extract_bm25_field_list (indexdef )
295+ aliases : dict [str , str ] = {}
296+ fields_ordered : list [str ] = []
297+ for raw in raw_fields :
298+ field_name = _extract_field_name (raw )
299+ if field_name is None :
300+ continue
301+ if field_name not in fields_ordered :
302+ fields_ordered .append (field_name )
303+ alias = _extract_alias (raw )
304+ if alias is not None :
305+ aliases [alias ] = field_name
199306
200307 output .append (
201308 IndexMeta (
202309 index_name = row .indexname ,
203310 key_field = key_field ,
204- fields = ( ),
205- aliases = {} ,
311+ fields = tuple ( fields_ordered ),
312+ aliases = aliases ,
206313 )
207314 )
208315 return output
0 commit comments