codekingpro/portable-devtools
115k
1import re2import sqlparse3from collections import namedtuple4from sqlparse.sql import Comparison, Identifier, Where5from .parseutils.utils import last_word, find_prev_keyword,\6 parse_partial_identifier7from .parseutils.tables import extract_tables8from .parseutils.ctes import isolate_query_ctes9 10 11Database = namedtuple("Database", [])12Schema = namedtuple("Schema", ["quoted"])13Schema.__new__.__defaults__ = (False,)14# FromClauseItem is a table/view/function used in the FROM clause15# `table_refs` contains the list of tables/... already in the statement,16# used to ensure that the alias we suggest is unique17FromClauseItem = namedtuple("FromClauseItem", "schema table_refs local_tables")18Table = namedtuple("Table", ["schema", "table_refs", "local_tables"])19TableFormat = namedtuple("TableFormat", [])20View = namedtuple("View", ["schema", "table_refs"])21# JoinConditions are suggested after ON, e.g. 'foo.barid = bar.barid'22JoinCondition = namedtuple("JoinCondition", ["table_refs", "parent"])23# Joins are suggested after JOIN, e.g. 'foo ON foo.barid = bar.barid'24Join = namedtuple("Join", ["table_refs", "schema"])25 26Function = namedtuple("Function", ["schema", "table_refs", "usage"])27# For convenience, don't require the `usage` argument in Function constructor28Function.__new__.__defaults__ = (None, tuple(), None)29Table.__new__.__defaults__ = (None, tuple(), tuple())30View.__new__.__defaults__ = (None, tuple())31FromClauseItem.__new__.__defaults__ = (None, tuple(), tuple())32 33Column = namedtuple(34 "Column",35 ["table_refs", "require_last_table", "local_tables", "qualifiable",36 "context"],37)38Column.__new__.__defaults__ = (None, None, tuple(), False, None)39 40Keyword = namedtuple("Keyword", ["last_token"])41Keyword.__new__.__defaults__ = (None,)42NamedQuery = namedtuple("NamedQuery", [])43Datatype = namedtuple("Datatype", ["schema"])44Alias = namedtuple("Alias", ["aliases"])45 46Path = namedtuple("Path", [])47 48 49class SqlStatement:50 def __init__(self, full_text, text_before_cursor):51 self.identifier = None52 self.word_before_cursor = word_before_cursor = last_word(53 text_before_cursor, include="many_punctuations"54 )55 full_text = _strip_named_query(full_text)56 text_before_cursor = _strip_named_query(text_before_cursor)57 58 full_text, text_before_cursor, self.local_tables = isolate_query_ctes(59 full_text, text_before_cursor60 )61 62 self.text_before_cursor_including_last_word = text_before_cursor63 64 # If we've partially typed a word then word_before_cursor won't be an65 # empty string. In that case we want to remove the partially typed66 # string before sending it to the sqlparser. Otherwise the last token67 # will always be the partially typed string which renders the smart68 # completion useless because it will always return the list of69 # keywords as completion.70 if self.word_before_cursor:71 if word_before_cursor[-1] == "(" or word_before_cursor[0] == "\\":72 parsed = sqlparse.parse(text_before_cursor)73 else:74 text_before_cursor = \75 text_before_cursor[: -len(word_before_cursor)]76 parsed = sqlparse.parse(text_before_cursor)77 self.identifier = parse_partial_identifier(word_before_cursor)78 else:79 parsed = sqlparse.parse(text_before_cursor)80 81 full_text, text_before_cursor, parsed = _split_multiple_statements(82 full_text, text_before_cursor, parsed83 )84 85 self.full_text = full_text86 self.text_before_cursor = text_before_cursor87 self.parsed = parsed88 89 self.last_token = parsed and \90 parsed.token_prev(len(parsed.tokens))[1] or ""91 92 def is_insert(self):93 return self.parsed.token_first().value.lower() == "insert"94 95 def get_tables(self, scope="full"):96 """Gets the tables available in the statement.97 param `scope:` possible values: 'full', 'insert', 'before'98 If 'insert', only the first table is returned.99 If 'before', only tables before the cursor are returned.100 If not 'insert' and the stmt is an insert, the first table is skipped.101 """102 tables = extract_tables(103 self.full_text if scope == "full" else self.text_before_cursor104 )105 if scope == "insert":106 tables = tables[:1]107 elif self.is_insert():108 tables = tables[1:]109 return tables110 111 def get_previous_token(self, token):112 return self.parsed.token_prev(self.parsed.token_index(token))[1]113 114 def get_identifier_schema(self):115 schema = \116 (self.identifier and self.identifier.get_parent_name()) or None117 # If schema name is unquoted, lower-case it118 if schema and self.identifier.value[0] != '"':119 schema = schema.lower()120 121 return schema122 123 def reduce_to_prev_keyword(self, n_skip=0):124 prev_keyword, self.text_before_cursor = find_prev_keyword(125 self.text_before_cursor, n_skip=n_skip126 )127 return prev_keyword128 129 130def suggest_type(full_text, text_before_cursor):131 """Takes the full_text that is typed so far and also the text before the132 cursor to suggest completion type and scope.133 134 Returns a tuple with a type of entity ('table', 'column' etc) and a scope.135 A scope for a column category will be a list of tables.136 """137 138 if full_text.startswith("\\i "):139 return (Path(),)140 141 # This is a temporary hack; the exception handling142 # here should be removed once sqlparse has been fixed143 try:144 stmt = SqlStatement(full_text, text_before_cursor)145 except (TypeError, AttributeError):146 return []147 148 return suggest_based_on_last_token(stmt.last_token, stmt)149 150 151named_query_regex = re.compile(r"^\s*\\ns\s+[A-z0-9\-_]+\s+")152 153 154def _strip_named_query(txt):155 """156 This will strip "save named query" command in the beginning of the line:157 '\ns zzz SELECT * FROM abc' -> 'SELECT * FROM abc'158 ' \ns zzz SELECT * FROM abc' -> 'SELECT * FROM abc'159 """160 161 if named_query_regex.match(txt):162 txt = named_query_regex.sub("", txt)163 return txt164 165 166function_body_pattern = re.compile(r"(\$.*?\$)([\s\S]*?)\1", re.M)167 168 169def _find_function_body(text):170 split = function_body_pattern.search(text)171 return (split.start(2), split.end(2)) if split else (None, None)172 173 174def _statement_from_function(full_text, text_before_cursor, statement):175 current_pos = len(text_before_cursor)176 body_start, body_end = _find_function_body(full_text)177 if body_start is None:178 return full_text, text_before_cursor, statement179 if not body_start <= current_pos < body_end:180 return full_text, text_before_cursor, statement181 full_text = full_text[body_start:body_end]182 text_before_cursor = text_before_cursor[body_start:]183 parsed = sqlparse.parse(text_before_cursor)184 return _split_multiple_statements(full_text, text_before_cursor, parsed)185 186 187def _split_multiple_statements(full_text, text_before_cursor, parsed):188 if len(parsed) > 1:189 # Multiple statements being edited -- isolate the current one by190 # cumulatively summing statement lengths to find the one that bounds191 # the current position192 current_pos = len(text_before_cursor)193 stmt_start, stmt_end = 0, 0194 195 for statement in parsed:196 stmt_len = len(str(statement))197 stmt_start, stmt_end = stmt_end, stmt_end + stmt_len198 199 if stmt_end >= current_pos:200 text_before_cursor = full_text[stmt_start:current_pos]201 full_text = full_text[stmt_start:]202 break203 204 elif parsed:205 # A single statement206 statement = parsed[0]207 else:208 # The empty string209 return full_text, text_before_cursor, None210 211 token2 = None212 if statement.get_type() in ("CREATE", "CREATE OR REPLACE"):213 token1 = statement.token_first()214 if token1:215 token1_idx = statement.token_index(token1)216 token2 = statement.token_next(token1_idx)[1]217 if token2 and token2.value.upper() == "FUNCTION":218 full_text, text_before_cursor, statement = _statement_from_function(219 full_text, text_before_cursor, statement220 )221 return full_text, text_before_cursor, statement222 223 224def suggest_based_on_last_token(token, stmt):225 226 if isinstance(token, str):227 token_v = token.lower()228 elif isinstance(token, Comparison):229 # If 'token' is a Comparison type such as230 # 'select * FROM abc a JOIN def d ON a.id = d.'. Then calling231 # token.value on the comparison type will only return the lhs of the232 # comparison. In this case a.id. So we need to do token.tokens to get233 # both sides of the comparison and pick the last token out of that234 # list.235 token_v = token.tokens[-1].value.lower()236 elif isinstance(token, Where):237 # sqlparse groups all tokens from the where clause into a single token238 # list. This means that token.value may be something like239 # 'where foo > 5 and '. We need to look "inside" token.tokens to handle240 # suggestions in complicated where clauses correctly241 prev_keyword = stmt.reduce_to_prev_keyword()242 return suggest_based_on_last_token(prev_keyword, stmt)243 elif isinstance(token, Identifier):244 # If the previous token is an identifier, we can suggest datatypes if245 # we're in a parenthesized column/field list, e.g.:246 # CREATE TABLE foo (Identifier <CURSOR>247 # CREATE FUNCTION foo (Identifier <CURSOR>248 # If we're not in a parenthesized list, the most likely scenario is the249 # user is about to specify an alias, e.g.:250 # SELECT Identifier <CURSOR>251 # SELECT foo FROM Identifier <CURSOR>252 prev_keyword, _ = find_prev_keyword(stmt.text_before_cursor)253 if prev_keyword and prev_keyword.value == "(":254 # Suggest datatypes255 return suggest_based_on_last_token("type", stmt)256 else:257 return (Keyword(),)258 else:259 token_v = token.value.lower()260 261 if not token:262 return (Keyword(),)263 elif token_v.endswith("("):264 p = sqlparse.parse(stmt.text_before_cursor)[0]265 266 if p.tokens and isinstance(p.tokens[-1], Where):267 # Four possibilities:268 # 1 - Parenthesized clause like "WHERE foo AND ("269 # Suggest columns/functions270 # 2 - Function call like "WHERE foo("271 # Suggest columns/functions272 # 3 - Subquery expression like "WHERE EXISTS ("273 # Suggest keywords, in order to do a subquery274 # 4 - Subquery OR array comparison like "WHERE foo = ANY("275 # Suggest columns/functions AND keywords. (If we wanted to276 # be really fancy, we could suggest only array-typed columns)277 278 column_suggestions = suggest_based_on_last_token("where", stmt)279 280 # Check for a subquery expression (cases 3 & 4)281 where = p.tokens[-1]282 prev_tok = where.token_prev(len(where.tokens) - 1)[1]283 284 if isinstance(prev_tok, Comparison):285 # e.g. "SELECT foo FROM bar WHERE foo = ANY("286 prev_tok = prev_tok.tokens[-1]287 288 prev_tok = prev_tok.value.lower()289 if prev_tok == "exists":290 return (Keyword(),)291 else:292 return column_suggestions293 294 # Get the token before the parens295 prev_tok = p.token_prev(len(p.tokens) - 1)[1]296 297 if (298 prev_tok and prev_tok.value and299 prev_tok.value.lower().split(" ")[-1] == "using"300 ):301 # tbl1 INNER JOIN tbl2 USING (col1, col2)302 tables = stmt.get_tables("before")303 304 # suggest columns that are present in more than one table305 return (306 Column(307 table_refs=tables,308 require_last_table=True,309 local_tables=stmt.local_tables,310 ),311 )312 313 elif p.token_first().value.lower() == "select":314 # If the lparen is preceeded by a space chances are we're about to315 # do a sub-select.316 if last_word(stmt.text_before_cursor,317 "all_punctuations").startswith("("):318 return (Keyword(),)319 prev_prev_tok = prev_tok and p.token_prev(p.token_index(prev_tok))[1]320 if prev_prev_tok and prev_prev_tok.normalized == "INTO":321 return (Column(table_refs=stmt.get_tables("insert"),322 context="insert"),)323 # We're probably in a function argument list324 return _suggest_expression(token_v, stmt)325 elif token_v == "set":326 return (Column(table_refs=stmt.get_tables(),327 local_tables=stmt.local_tables),)328 elif token_v in ("select", "where", "having", "order by", "distinct"):329 return _suggest_expression(token_v, stmt)330 elif token_v == "as":331 # Don't suggest anything for aliases332 return ()333 elif (token_v.endswith("join") and token.is_keyword) or (334 token_v in ("copy", "from", "update", "into", "describe", "truncate")335 ):336 337 schema = stmt.get_identifier_schema()338 tables = extract_tables(stmt.text_before_cursor)339 is_join = token_v.endswith("join") and token.is_keyword340 341 # Suggest tables from either the currently-selected schema or the342 # public schema if no schema has been specified343 suggest = []344 345 if not schema:346 # Suggest schemas347 suggest.insert(0, Schema())348 349 if token_v == "from" or is_join:350 suggest.append(351 FromClauseItem(352 schema=schema, table_refs=tables,353 local_tables=stmt.local_tables354 )355 )356 elif token_v == "truncate":357 suggest.append(Table(schema))358 else:359 suggest.extend((Table(schema), View(schema)))360 361 if is_join and _allow_join(stmt.parsed):362 tables = stmt.get_tables("before")363 suggest.append(Join(table_refs=tables, schema=schema))364 365 return tuple(suggest)366 367 elif token_v == "function":368 schema = stmt.get_identifier_schema()369 370 # stmt.get_previous_token will fail for e.g.371 # `SELECT 1 FROM functions WHERE function:`372 try:373 prev = stmt.get_previous_token(token).value.lower()374 if prev in ("drop", "alter", "create", "create or replace"):375 376 # Suggest functions from either the currently-selected schema377 # or the public schema if no schema has been specified378 suggest = []379 380 if not schema:381 # Suggest schemas382 suggest.insert(0, Schema())383 384 suggest.append(Function(schema=schema, usage="signature"))385 return tuple(suggest)386 387 except ValueError:388 pass389 return tuple()390 391 elif token_v in ("table", "view"):392 # E.g. 'ALTER TABLE <tablname>'393 rel_type = \394 {"table": Table, "view": View, "function": Function}[token_v]395 schema = stmt.get_identifier_schema()396 if schema:397 return (rel_type(schema=schema),)398 else:399 return (Schema(), rel_type(schema=schema))400 401 elif token_v == "column":402 # E.g. 'ALTER TABLE foo ALTER COLUMN bar403 return (Column(table_refs=stmt.get_tables()),)404 405 elif token_v == "on":406 tables = stmt.get_tables("before")407 parent = \408 (stmt.identifier and stmt.identifier.get_parent_name()) or None409 if parent:410 # "ON parent.<suggestion>"411 # parent can be either a schema name or table alias412 filteredtables = tuple(t for t in tables if identifies(parent, t))413 sugs = [414 Column(table_refs=filteredtables,415 local_tables=stmt.local_tables),416 Table(schema=parent),417 View(schema=parent),418 Function(schema=parent),419 ]420 if filteredtables and _allow_join_condition(stmt.parsed):421 sugs.append(JoinCondition(table_refs=tables,422 parent=filteredtables[-1]))423 return tuple(sugs)424 else:425 # ON <suggestion>426 # Use table alias if there is one, otherwise the table name427 aliases = tuple(t.ref for t in tables)428 if _allow_join_condition(stmt.parsed):429 return (430 Alias(aliases=aliases),431 JoinCondition(table_refs=tables, parent=None),432 )433 else:434 return (Alias(aliases=aliases),)435 436 elif token_v in ("c", "use", "database", "template"):437 # "\c <db", "use <db>", "DROP DATABASE <db>",438 # "CREATE DATABASE <newdb> WITH TEMPLATE <db>"439 return (Database(),)440 elif token_v == "schema":441 # DROP SCHEMA schema_name, SET SCHEMA schema name442 prev_keyword = stmt.reduce_to_prev_keyword(n_skip=2)443 quoted = prev_keyword and prev_keyword.value.lower() == "set"444 return (Schema(quoted),)445 elif token_v.endswith(",") or token_v in ("=", "and", "or"):446 prev_keyword = stmt.reduce_to_prev_keyword()447 if prev_keyword:448 return suggest_based_on_last_token(prev_keyword, stmt)449 else:450 return ()451 elif token_v in ("type", "::"):452 # ALTER TABLE foo SET DATA TYPE bar453 # SELECT foo::bar454 # Note that tables are a form of composite type in postgresql, so455 # they're suggested here as well456 schema = stmt.get_identifier_schema()457 suggestions = [Datatype(schema=schema), Table(schema=schema)]458 if not schema:459 suggestions.append(Schema())460 return tuple(suggestions)461 elif token_v in {"alter", "create", "drop"}:462 return (Keyword(token_v.upper()),)463 elif token.is_keyword:464 # token is a keyword we haven't implemented any special handling for465 # go backwards in the query until we find one we do recognize466 prev_keyword = stmt.reduce_to_prev_keyword(n_skip=1)467 if prev_keyword:468 return suggest_based_on_last_token(prev_keyword, stmt)469 else:470 return (Keyword(token_v.upper()),)471 else:472 return (Keyword(),)473 474 475def _suggest_expression(token_v, stmt):476 """477 Return suggestions for an expression, taking account of any partially-typed478 identifier's parent, which may be a table alias or schema name.479 """480 parent = stmt.identifier.get_parent_name() if stmt.identifier else []481 tables = stmt.get_tables()482 483 if parent:484 tables = tuple(t for t in tables if identifies(parent, t))485 return (486 Column(table_refs=tables, local_tables=stmt.local_tables),487 Table(schema=parent),488 View(schema=parent),489 Function(schema=parent),490 )491 492 return (493 Column(table_refs=tables, local_tables=stmt.local_tables,494 qualifiable=True),495 Function(schema=None),496 Keyword(token_v.upper()),497 )498 499 500def identifies(id, ref):501 """Returns true if string `id` matches TableReference `ref`"""502 503 return (504 id == ref.alias or id == ref.name or505 (ref.schema and (id == ref.schema + "." + ref.name))506 )507 508 509def _allow_join_condition(statement):510 """511 Tests if a join condition should be suggested512 513 We need this to avoid bad suggestions when entering e.g.514 select * from tbl1 a join tbl2 b on a.id = <cursor>515 So check that the preceding token is a ON, AND, or OR keyword, instead of516 e.g. an equals sign.517 518 :param statement: an sqlparse.sql.Statement519 :return: boolean520 """521 522 if not statement or not statement.tokens:523 return False524 525 last_tok = statement.token_prev(len(statement.tokens))[1]526 return last_tok.value.lower() in ("on", "and", "or")527 528 529def _allow_join(statement):530 """531 Tests if a join should be suggested532 533 We need this to avoid bad suggestions when entering e.g.534 select * from tbl1 a join tbl2 b <cursor>535 So check that the preceding token is a JOIN keyword536 537 :param statement: an sqlparse.sql.Statement538 :return: boolean539 """540 541 if not statement or not statement.tokens:542 return False543 544 last_tok = statement.token_prev(len(statement.tokens))[1]545 return last_tok.value.lower().endswith("join") and \546 last_tok.value.lower() not in ("cross join", "natural join",)547 