diff --git a/patroni/ctl.py b/patroni/ctl.py index a8ca193c..4f2fcd0b 100644 --- a/patroni/ctl.py +++ b/patroni/ctl.py @@ -31,7 +31,7 @@ from patroni.postgresql.misc import postgres_version_to_int from patroni.utils import cluster_as_json, patch_config, polling_loop from patroni.request import PatroniRequest from patroni.version import __version__ -from prettytable import PrettyTable +from prettytable import ALL, FRAME, PrettyTable from six.moves.urllib_parse import urlparse CONFIG_DIR_PATH = click.get_app_dir('patroni') @@ -46,6 +46,34 @@ class PatroniCtlException(ClickException): pass +class PatronictlPrettyTable(PrettyTable): + + def __init__(self, header, *args, **kwargs): + PrettyTable.__init__(self, *args, **kwargs) + self.__table_header = header + self.__hline_num = 0 + self.__hline = None + + def _is_first_hline(self): + return self.__hline_num == 0 + + def _set_hline(self, value): + self.__hline = value + + def _get_hline(self): + ret = self.__hline + + # Inject nice table header + if self._is_first_hline() and self.__table_header: + header = self.__table_header[:len(ret) - 2] + ret = "".join([ret[0], header, ret[1 + len(header):]]) + + self.__hline_num += 1 + return ret + + _hrule = property(_get_hline, _set_hline) + + def parse_dcs(dcs): if dcs is None: return None @@ -94,7 +122,7 @@ def store_config(config, path): yaml.dump(config, fd) -option_format = click.option('--format', '-f', 'fmt', help='Output format (pretty, json, yaml)', default='pretty') +option_format = click.option('--format', '-f', 'fmt', help='Output format (pretty, tsv, json, yaml)', default='pretty') option_watchrefresh = click.option('-w', '--watch', type=float, help='Auto update the screen every X seconds') option_watch = click.option('-W', is_flag=True, help='Auto update the screen every 2 seconds') option_force = click.option('--force', is_flag=True, help='Do not ask for confirmation at any point') @@ -137,31 +165,33 @@ def request_patroni(member, method='GET', endpoint=None, data=None): return request_executor(member, method, endpoint, data) -def print_output(columns, rows=None, alignment=None, fmt='pretty', header=True, delimiter='\t'): - rows = rows or [] - if fmt == 'pretty': - t = PrettyTable(columns) - for k, v in (alignment or {}).items(): - t.align[k] = v - for r in rows: - t.add_row(r) - click.echo(t) - return +def print_output(columns, rows, alignment=None, fmt='pretty', header=None, delimiter='\t'): + if fmt in {'json', 'yaml', 'yml'}: + elements = [{k: v for k, v in zip(columns, r) if not header or str(v)} for r in rows] + func = json.dumps if fmt == 'json' else format_config_for_editing + click.echo(func(elements)) + elif fmt in {'pretty', 'tsv'}: + list_cluster = bool(header and columns and columns[0] == 'Cluster') + if list_cluster and 'Tags' in columns: # we want to format member tags as YAML + i = columns.index('Tags') + for row in rows: + if row[i]: + row[i] = format_config_for_editing(row[i], fmt == 'tsv').strip() + if list_cluster and fmt == 'pretty': # skip cluster name if pretty-printing + columns = columns[1:] if columns else [] + rows = [row[1:] for row in rows] - if fmt in ['json', 'yaml', 'yml']: - elements = [dict(zip(columns, r)) for r in rows] - if fmt == 'json': - click.echo(json.dumps(elements)) - elif fmt in ('yaml', 'yml'): - click.echo(yaml.safe_dump(elements, encoding=None, default_flow_style=False, allow_unicode=True, width=200)) - - if fmt == 'tsv': - if columns is not None and header: - click.echo(delimiter.join(columns)) - - for r in rows: - c = [str(c) for c in r] - click.echo(delimiter.join(c)) + if fmt == 'tsv': + for r in ([columns] if columns else []) + rows: + click.echo(delimiter.join(map(str, r))) + else: + hrules = ALL if any(any(isinstance(c, six.string_types) and '\n' in c for c in r) for r in rows) else FRAME + table = PatronictlPrettyTable(header, columns, hrules=hrules) + for k, v in (alignment or {}).items(): + table.align[k] = v + for r in rows: + table.add_row(r) + click.echo(table) def watching(w, watch, max_count=None, clear=True): @@ -310,8 +340,7 @@ def dsn(obj, cluster_name, role, member): @ctl.command('query', help='Query a Patroni PostgreSQL member') @arg_cluster_name -@option_format -@click.option('--format', 'fmt', help='Output format (pretty, json)', default='tsv') +@click.option('--format', 'fmt', help='Output format (pretty, tsv, json, yaml)', default='tsv') @click.option('--file', '-f', 'p_file', help='Execute the SQL commands from this file', type=click.File('rb')) @click.option('--password', help='force password prompt', is_flag=True) @click.option('-U', '--username', help='database user name', type=str) @@ -368,8 +397,8 @@ def query( if cursor is None: cluster = dcs.get_cluster() - output, cursor = query_member(cluster, cursor, member, role, command, connect_parameters) - print_output(None, output, fmt=fmt, delimiter=delimiter) + output, header = query_member(cluster, cursor, member, role, command, connect_parameters) + print_output(header, output, fmt=fmt, delimiter=delimiter) def query_member(cluster, cursor, member, role, command, connect_parameters): @@ -386,15 +415,8 @@ def query_member(cluster, cursor, member, role, command, connect_parameters): logging.debug(message) return [[timestamp(0), message]], None - cursor.execute('SELECT pg_catalog.pg_is_in_recovery()') - in_recovery = cursor.fetchone()[0] - - if in_recovery and role == 'master' or not in_recovery and role == 'replica': - cursor.connection.close() - return None, None - cursor.execute(command) - return cursor.fetchall(), cursor + return cursor.fetchall(), [d.name for d in cursor.description] except (psycopg2.OperationalError, psycopg2.DatabaseError) as oe: logging.debug(oe) if cursor is not None and not cursor.connection.closed: @@ -727,6 +749,7 @@ def switchover(obj, cluster_name, master, candidate, force, scheduled): def output_members(cluster, name, extended=False, fmt='pretty'): rows = [] logging.debug(cluster) + initialize = {None: 'uninitialized', '': 'initializing'}.get(cluster.initialize, cluster.initialize) cluster = cluster_as_json(cluster) columns = ['Cluster', 'Member', 'Host', 'Role', 'State', 'TL', 'Lag in MB'] @@ -746,8 +769,7 @@ def output_members(cluster, name, extended=False, fmt='pretty'): m.update(cluster=name, member=m['name'], host=m.get('host'), tl=m.get('timeline', ''), role='' if m['role'] == 'replica' else m['role'].replace('_', ' ').title(), lag_in_mb=round(lag/1024/1024) if isinstance(lag, six.integer_types) else lag, - pending_restart='*' if m.get('pending_restart') else '', - tags=json.dumps(m['tags']) if m.get('tags') else '') + pending_restart='*' if m.get('pending_restart') else '') if append_port and m['host'] and m.get('port'): m['host'] = ':'.join([m['host'], str(m['port'])]) @@ -760,7 +782,8 @@ def output_members(cluster, name, extended=False, fmt='pretty'): rows.append([m.get(n.lower().replace(' ', '_'), '') for n in columns]) - print_output(columns, rows, {'Lag in MB': 'r', 'TL': 'r', 'Tags': 'l'}, fmt) + print_output(columns, rows, {'Lag in MB': 'r', 'TL': 'r', 'Tags': 'l'}, + fmt, ' Cluster: {0} ({1}) '.format(name, initialize)) if fmt != 'pretty': # Omit service info when using machine-readable formats return @@ -1006,12 +1029,12 @@ def show_diff(before_editing, after_editing): click.echo(line.rstrip('\n')) -def format_config_for_editing(data): +def format_config_for_editing(data, default_flow_style=False): """Formats configuration as YAML for human consumption. :param data: configuration as nested dictionaries :returns unicode YAML of the configuration""" - return yaml.safe_dump(data, default_flow_style=False, encoding=None, allow_unicode=True) + return yaml.safe_dump(data, default_flow_style=default_flow_style, encoding=None, allow_unicode=True, width=200) def apply_config_changes(before_editing, data, kvpairs): diff --git a/tests/__init__.py b/tests/__init__.py index 7adea2a0..8a075d61 100644 --- a/tests/__init__.py +++ b/tests/__init__.py @@ -76,6 +76,7 @@ class MockCursor(object): self.closed = False self.rowcount = 0 self.results = [] + self.description = [Mock()] def execute(self, sql, *params): if sql.startswith('blabla'): diff --git a/tests/test_ctl.py b/tests/test_ctl.py index eb2d9793..b72785b6 100644 --- a/tests/test_ctl.py +++ b/tests/test_ctl.py @@ -206,9 +206,6 @@ class TestCtl(unittest.TestCase): rows = query_member(None, None, None, 'master', 'SELECT pg_catalog.pg_is_in_recovery()', {}) self.assertTrue('False' in str(rows)) - rows = query_member(None, None, None, 'replica', 'SELECT pg_catalog.pg_is_in_recovery()', {}) - self.assertEqual(rows, (None, None)) - with patch.object(MockCursor, 'execute', Mock(side_effect=OperationalError('bla'))): rows = query_member(None, None, None, 'replica', 'SELECT pg_catalog.pg_is_in_recovery()', {})