Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes15kdownloads
sqlcompletion.py547 linesDownload Raw Back to sqlautocomplete
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 
codekingpro/portable-devtools · Team Ai