diff --git a/pgcli/main.py b/pgcli/main.py index e0bc2dce3..29f5311c4 100755 --- a/pgcli/main.py +++ b/pgcli/main.py @@ -7,6 +7,7 @@ import traceback import logging from time import time +from codecs import open import click import sqlparse @@ -97,6 +98,8 @@ def register_special_commands(self): 'Refresh auto-completions.', arg_type=NO_QUERY) self.pgspecial.register(self.refresh_completions, '\\refresh', '\\refresh', 'Refresh auto-completions.', arg_type=NO_QUERY) + self.pgspecial.register(self.execute_from_file, '\\i', '\\i filename', + 'Execute commands from file.') def change_db(self, pattern, **_): if pattern: @@ -108,6 +111,18 @@ def change_db(self, pattern, **_): yield (None, None, None, 'You are now connected to database "%s" as ' 'user "%s"' % (self.pgexecute.dbname, self.pgexecute.user)) + def execute_from_file(self, pattern, **_): + if not pattern: + message = '\\i: missing required argument' + return [(None, None, None, message)] + try: + with open(os.path.expanduser(pattern), encoding='utf-8') as f: + query = f.read() + except IOError as e: + return [(None, None, None, str(e))] + + return self.pgexecute.run(query, self.pgspecial) + def initialize_logging(self): log_file = self.config['main']['log_file'] diff --git a/pgcli/packages/pgspecial/iocommands.py b/pgcli/packages/pgspecial/iocommands.py index 0a2969693..621b5862e 100644 --- a/pgcli/packages/pgspecial/iocommands.py +++ b/pgcli/packages/pgspecial/iocommands.py @@ -1,11 +1,9 @@ import re import logging -from codecs import open -from os.path import expanduser import click from .namedqueries import namedqueries -from .main import special_command, NO_QUERY from . import export +from .main import special_command _logger = logging.getLogger(__name__) @@ -68,30 +66,6 @@ def open_external_editor(filename=None, sql=''): return (query, message) -@special_command('\\i', '\\i file', 'Execute commands from file.') -def execute_from_file(cur, pattern, **_): - if pattern: - try: - query = read_from_file(pattern) - except IOError as e: - message = 'Error reading file: %s' % pattern - message = message + ' Error was: ' + str(e) - return [(None, None, None, message)] - else: - message = '\\i: missing required argument' - return [(None, None, None, message)] - cur.execute(query) - if cur.description: - headers = [x[0] for x in cur.description] - return [(None, cur, headers, cur.statusmessage)] - else: - return [(None, None, None, cur.statusmessage)] - -def read_from_file(path): - with open(expanduser(path), encoding='utf-8') as f: - contents = f.read() - return contents - @special_command('\\n', '\\n[+] [name]', 'List or execute named queries.') def execute_named_query(cur, pattern, **_): """Returns (title, rows, headers, status)"""