codekingpro/portable-devtools
114k
1from sqlparse import parse2from sqlparse.tokens import Keyword, CTE, DML3from sqlparse.sql import Identifier, IdentifierList, Parenthesis4from collections import namedtuple5from .meta import TableMetadata, ColumnMetadata6 7 8# TableExpression is a namedtuple representing a CTE, used internally9# name: cte alias assigned in the query10# columns: list of column names11# start: index into the original string of the left parens starting the CTE12# stop: index into the original string of the right parens ending the CTE13TableExpression = namedtuple("TableExpression", "name columns start stop")14 15 16def isolate_query_ctes(full_text, text_before_cursor):17 """Simplify a query by converting CTEs into table metadata objects"""18 19 if not full_text or not full_text.strip():20 return full_text, text_before_cursor, tuple()21 22 ctes, _ = extract_ctes(full_text)23 if not ctes:24 return full_text, text_before_cursor, ()25 26 current_position = len(text_before_cursor)27 meta = []28 29 for cte in ctes:30 if cte.start < current_position < cte.stop:31 # Currently editing a cte - treat its body as the current full_text32 text_before_cursor = full_text[cte.start: current_position]33 full_text = full_text[cte.start: cte.stop]34 return full_text, text_before_cursor, meta35 36 # Append this cte to the list of available table metadata37 cols = (ColumnMetadata(name, None, ()) for name in cte.columns)38 meta.append(TableMetadata(cte.name, cols))39 40 # Editing past the last cte (ie the main body of the query)41 full_text = full_text[ctes[-1].stop:]42 text_before_cursor = text_before_cursor[ctes[-1].stop: current_position]43 44 return full_text, text_before_cursor, tuple(meta)45 46 47def extract_ctes(sql):48 """Extract constant table expresseions from a query49 50 Returns tuple (ctes, remainder_sql)51 52 ctes is a list of TableExpression namedtuples53 remainder_sql is the text from the original query after the CTEs have54 been stripped.55 """56 57 p = parse(sql)[0]58 59 # Make sure the first meaningful token is "WITH" which is necessary to60 # define CTEs61 idx, tok = p.token_next(-1, skip_ws=True, skip_cm=True)62 if not (tok and tok.ttype == CTE):63 return [], sql64 65 # Get the next (meaningful) token, which should be the first CTE66 idx, tok = p.token_next(idx)67 if not tok:68 return ([], "")69 start_pos = token_start_pos(p.tokens, idx)70 ctes = []71 72 if isinstance(tok, IdentifierList):73 # Multiple ctes74 for t in tok.get_identifiers():75 cte_start_offset = token_start_pos(tok.tokens, tok.token_index(t))76 cte = get_cte_from_token(t, start_pos + cte_start_offset)77 if not cte:78 continue79 ctes.append(cte)80 elif isinstance(tok, Identifier):81 # A single CTE82 cte = get_cte_from_token(tok, start_pos)83 if cte:84 ctes.append(cte)85 86 idx = p.token_index(tok) + 187 88 # Collapse everything after the ctes into a remainder query89 remainder = "".join(str(tok) for tok in p.tokens[idx:])90 91 return ctes, remainder92 93 94def get_cte_from_token(tok, pos0):95 cte_name = tok.get_real_name()96 if not cte_name:97 return None98 99 # Find the start position of the opening parens enclosing the cte body100 idx, parens = tok.token_next_by(Parenthesis)101 if not parens:102 return None103 104 start_pos = pos0 + token_start_pos(tok.tokens, idx)105 cte_len = len(str(parens)) # includes parens106 stop_pos = start_pos + cte_len107 108 column_names = extract_column_names(parens)109 110 return TableExpression(cte_name, column_names, start_pos, stop_pos)111 112 113def extract_column_names(parsed):114 # Find the first DML token to check if it's a115 # SELECT or INSERT/UPDATE/DELETE116 idx, tok = parsed.token_next_by(t=DML)117 tok_val = tok and tok.value.lower()118 119 if tok_val in ("insert", "update", "delete"):120 # Jump ahead to the RETURNING clause where the list of column names is121 idx, tok = parsed.token_next_by(idx, (Keyword, "returning"))122 elif tok_val != "select":123 # Must be invalid CTE124 return ()125 126 # The next token should be either a column name, or a list of column names127 idx, tok = parsed.token_next(idx, skip_ws=True, skip_cm=True)128 return tuple(t.get_name() for t in _identifiers(tok))129 130 131def token_start_pos(tokens, idx):132 return sum(len(str(t)) for t in tokens[:idx])133 134 135def _identifiers(tok):136 if isinstance(tok, IdentifierList):137 for t in tok.get_identifiers():138 # NB: IdentifierList.get_identifiers() can return non-identifiers!139 if isinstance(t, Identifier):140 yield t141 elif isinstance(tok, Identifier):142 yield tok143 