Clean up obsolete functions and simplify cql_version handling in cqlsh

Patch by Brad Schoening; reviewed by brandonwilliams and edimitrova for
CASSANDRA-18787
This commit is contained in:
Brad Schoening 2023-08-30 23:07:48 -04:00 committed by Brandon Williams
parent 5aa7da5ebd
commit abe09cff34
6 changed files with 40 additions and 98 deletions

View File

@ -1,4 +1,5 @@
5.1
* Clean up obsolete functions and simplify cql_version handling in cqlsh (CASSANDRA-18787)
Merged from 5.0:
* Fix Depends and Build-Depends for Java for Debian packages (CASSANDRA-18809)
* Update command line flags --add-exports and --add-opens for JDK17 (CASSANDRA-18439)

View File

@ -99,18 +99,11 @@ port = 9042
[csv]
;; The size limit for parsed fields
; field_size_limit = 131072
[tracing]
;; The max number of seconds to wait for a trace to complete
; max_trace_wait = 10.0
;[ssl]
; certfile = ~/keys/cassandra.cert

View File

@ -17,7 +17,6 @@
import cmd
import codecs
import configparser
import csv
import getpass
import argparse
import os
@ -277,22 +276,10 @@ class DecodeError(Exception):
return '<%s %s>' % (self.__class__.__name__, self.message())
def maybe_ensure_text(val):
return str(val) if val else val
class FormatError(DecodeError):
verb = 'format'
def full_cql_version(ver):
while ver.count('.') < 2:
ver += '.0'
ver_parts = ver.split('-', 1) + ['']
vertuple = tuple(list(map(int, ver_parts[0].split('.'))) + [ver_parts[1]])
return ver, vertuple
def format_value(val, cqltype, encoding, addcolor=False, date_time_format=None,
float_precision=None, colormap=None, nullval=None):
if isinstance(val, DecodeError):
@ -412,7 +399,7 @@ class Shell(cmd.Cmd):
self.session.row_factory = ordered_dict_factory
self.session.default_consistency_level = cassandra.ConsistencyLevel.ONE
self.get_connection_versions()
self.set_expanded_cql_version(self.connection_versions['cql'])
self.set_cql_version(self.connection_versions['cql'])
self.current_keyspace = keyspace
@ -450,6 +437,13 @@ class Shell(cmd.Cmd):
self.single_statement = single_statement
self.is_subshell = is_subshell
self.cql_version = None
self.cql_version_str = None
# configure length of history shown
self.max_history_length_shown = 50
self.lastcmd = ""
def check_build_versions(self):
baseversion = self.connection_versions['build']
extra = baseversion.rfind('-')
@ -462,13 +456,12 @@ class Shell(cmd.Cmd):
def batch_mode(self):
return not self.tty
def set_expanded_cql_version(self, ver):
ver, vertuple = full_cql_version(ver)
self.cql_version = ver
self.cql_ver_tuple = vertuple
def cqlver_atleast(self, major, minor=0, patch=0):
return self.cql_ver_tuple[:3] >= (major, minor, patch)
def set_cql_version(self, ver):
v = list(map(int, (ver.split("."))))
while (len(v) < 3):
v.append(0)
self.cql_version = tuple(v)
self.cql_version_str = ".".join(map(str, v))
def myformat_value(self, val, cqltype=None, **kwargs):
if isinstance(val, DecodeError):
@ -510,12 +503,8 @@ class Shell(cmd.Cmd):
self.port))
def show_version(self):
vers = self.connection_versions.copy()
vers['shver'] = version
# system.Versions['cql'] apparently does not reflect changes with
# set_cql_version.
vers['cql'] = self.cql_version
print("[cqlsh %(shver)s | Cassandra %(build)s | CQL spec %(cql)s | Native protocol v%(protocol)s]" % vers)
vers = self.connection_versions
print(f"[cqlsh {version} | Cassandra {vers['build']} | CQL spec {self.cql_version_str} | Native protocol v{vers['protocol']}]")
def show_session(self, sessionid, partial_session=False):
print_trace_session(self, self.session, sessionid, partial_session)
@ -1321,7 +1310,7 @@ class Shell(cmd.Cmd):
def describe_list(self, rows):
"""
Print the output for all the DESCRIBE queries for element names (e.g DESCRIBE TABLES, DESCRIBE FUNCTIONS ...)
Print the output for all the DESCRIBE queries for element names (e.g. DESCRIBE TABLES, DESCRIBE FUNCTIONS ...)
"""
keyspace = None
names = list()
@ -1348,7 +1337,7 @@ class Shell(cmd.Cmd):
def describe_element(self, rows):
"""
Print the output for all the DESCRIBE queries where an element name as been specified (e.g DESCRIBE TABLE, DESCRIBE INDEX ...)
Print the output for all the DESCRIBE queries where an element name as been specified (e.g. DESCRIBE TABLE, DESCRIBE INDEX ...)
"""
for row in rows:
print('')
@ -1359,7 +1348,7 @@ class Shell(cmd.Cmd):
"""
Print the output for a DESCRIBE CLUSTER query.
If a specified keyspace was in use the returned ResultSet will contains a 'range_ownership' column,
If a specified keyspace was in use the returned ResultSet will contain a 'range_ownership' column,
otherwise not.
"""
for row in rows:
@ -1957,9 +1946,6 @@ class Shell(cmd.Cmd):
delims += '.'
readline.set_completer_delims(delims)
# configure length of history shown
self.max_history_length_shown = 50
def save_history(self):
if readline is not None:
try:
@ -2057,8 +2043,7 @@ def should_use_color():
if int(stdout.strip()) < 8:
return False
except (OSError, ImportError, ValueError):
# oh well, we tried. at least we know there's a $TERM and it's
# not "dumb".
# at least it's a $TERM, and it's not "dumb".
pass
return True
@ -2100,7 +2085,6 @@ def read_options(cmdlineargs, environment=os.environ):
DEFAULT_FLOAT_PRECISION)
argvalues.double_precision = option_with_default(configs.getint, 'ui', 'double_precision',
DEFAULT_DOUBLE_PRECISION)
argvalues.field_size_limit = option_with_default(configs.getint, 'csv', 'field_size_limit', csv.field_size_limit())
argvalues.max_trace_wait = option_with_default(configs.getfloat, 'tracing', 'max_trace_wait',
DEFAULT_MAX_TRACE_WAIT)
argvalues.timezone = option_with_default(configs.get, 'ui', 'timezone', None)
@ -2166,15 +2150,12 @@ def read_options(cmdlineargs, environment=os.environ):
print("\nWarning: Using a password on the command line interface can be insecure."
"\nRecommendation: use the credentials file to securely provide the password.\n", file=sys.stderr)
# Make sure some user values read from the command line are in unicode
options.execute = maybe_ensure_text(options.execute)
options.username = maybe_ensure_text(options.username)
options.password = maybe_ensure_text(options.password)
options.keyspace = maybe_ensure_text(options.keyspace)
hostname = option_with_default(configs.get, 'connection', 'hostname', DEFAULT_HOST)
port = option_with_default(configs.get, 'connection', 'port', DEFAULT_PORT)
hostname = environment.get('CQLSH_HOST', hostname)
port = environment.get('CQLSH_PORT', port)
try:
options.connect_timeout = int(options.connect_timeout)
except ValueError:
@ -2187,9 +2168,6 @@ def read_options(cmdlineargs, environment=os.environ):
parser.error('"%s" is not a valid request timeout.' % (options.request_timeout,))
options.request_timeout = DEFAULT_REQUEST_TIMEOUT_SECONDS
hostname = environment.get('CQLSH_HOST', hostname)
port = environment.get('CQLSH_PORT', port)
if len(arguments) > 0:
hostname = arguments[0]
if len(arguments) > 1:
@ -2209,12 +2187,6 @@ def read_options(cmdlineargs, environment=os.environ):
else:
options.color = should_use_color()
if options.cqlversion is not None:
options.cqlversion, cqlvertup = full_cql_version(options.cqlversion)
if cqlvertup[0] < 3:
parser.error('%r is not a supported CQL version.' % options.cqlversion)
options.cqlmodule = cql3handling
try:
port = int(port)
except ValueError:
@ -2222,9 +2194,9 @@ def read_options(cmdlineargs, environment=os.environ):
return options, hostname, port
def setup_cqlruleset(cqlmodule):
def setup_cqlruleset():
global cqlruleset
cqlruleset = cqlmodule.CqlRuleSet
cqlruleset = cql3handling.CqlRuleSet
cqlruleset.append_rules(cqlshhandling.cqlsh_extra_syntax_rules)
for rulename, termname, func in cqlshhandling.cqlsh_syntax_completers:
cqlruleset.completer_for(rulename, termname)(func)
@ -2253,7 +2225,7 @@ def insert_driver_hooks():
pass
# Display milliseconds when datetime overflows (CASSANDRA-10625), E.g., the year 10000.
# Native datetime types blow up outside of datetime.[MIN|MAX]_YEAR. We will fall back to an int timestamp
# Native datetime types blow up outside datetime.[MIN|MAX]_YEAR. We will fall back to an int timestamp
def deserialize_date_fallback_int(byts, protocol_version):
timestamp_ms = int64_unpack(byts)
try:
@ -2277,8 +2249,7 @@ def main(cmdline, pkgpath):
(options, hostname, port) = read_options(cmdline)
setup_docspath(pkgpath)
setup_cqlruleset(options.cqlmodule)
csv.field_size_limit(options.field_size_limit)
setup_cqlruleset()
if options.file is None:
stdin = None
@ -2331,7 +2302,7 @@ def main(cmdline, pkgpath):
except ImportError:
sys.stderr.write("Warning: Timezone defined and 'pytz' module for timezone conversion not installed. Timestamps will be displayed in UTC timezone.\n\n")
# try auto-detect timezone if tzlocal is installed
# try to auto-detect timezone if tzlocal is installed
if not timezone:
try:
from tzlocal import get_localzone

View File

@ -27,10 +27,9 @@ from collections import defaultdict
from cassandra.cqltypes import EMPTY
from cassandra.util import datetime_from_timestamp
from .displaying import colorme, get_str, FormattedValue, DEFAULT_VALUE_COLORS, NO_COLOR_MAP
from .util import UTC
unicode_controlchars_re = re.compile(r'[\x00-\x1f\x7f-\xa0]')
controlchars_re = re.compile(r'[\x00-\x1f\x7f-\xff]')
UNICODE_CONTROLCHARS_RE = re.compile(r'[\x00-\x1f\x7f-\xa0]')
CONTROLCHARS_RE = re.compile(r'[\x00-\x1f\x7f-\xff]')
def _show_control_chars(match):
@ -84,7 +83,7 @@ def format_by_type(val, cqltype, encoding, colormap=None, addcolor=False,
def color_text(bval, colormap, displaywidth=None):
# note that here, we render natural backslashes as just backslashes,
# in the same color as surrounding text, when using color. When not
# using color, we need to double up the backslashes so it's not
# using color, we need to double up the backslashes, so it's not
# ambiguous. This introduces the unique difficulty of having different
# display widths for the colored and non-colored versions. To avoid
# adding the smarts to handle that in to FormattedValue, we just
@ -121,7 +120,7 @@ class DateTimeFormat:
class CqlType:
"""
A class for converting a string into a cql type name that can match a formatter
and a list of its sub-types, if any.
and a list of its subtypes, if any.
"""
pattern = re.compile('^([^<]*)<(.*)>$') # *<*>
@ -135,8 +134,8 @@ class CqlType:
def get_n_sub_types(self, num):
"""
Return the sub-types if the requested number matches the length of the sub-types (tuples)
or the first sub-type times the number requested if the length of the sub-types is one (list, set),
Return the subtypes if the requested number matches the length of the subtypes (tuples)
or the first subtype times the number requested if the length of the subtypes is one (list, set),
otherwise raise an exception
"""
if len(self.sub_types) == num:
@ -202,7 +201,7 @@ class CqlType:
def format_value_default(val, colormap, **_):
val = str(val)
escapedval = val.replace('\\', '\\\\')
bval = controlchars_re.sub(_show_control_chars, escapedval)
bval = CONTROLCHARS_RE.sub(_show_control_chars, escapedval)
return bval if colormap is NO_COLOR_MAP else color_text(bval, colormap)
@ -359,7 +358,7 @@ formatter_for('timestamp')(format_value_timestamp)
def strftime(time_format, seconds, microseconds=0, timezone=None):
ret_dt = datetime_from_timestamp(seconds) + datetime.timedelta(microseconds=microseconds)
ret_dt = ret_dt.replace(tzinfo=UTC())
ret_dt = ret_dt.replace(tzinfo=datetime.timezone.utc)
if timezone:
ret_dt = ret_dt.astimezone(timezone)
try:
@ -475,7 +474,7 @@ def format_value_text(val, encoding, colormap, quote=False, **_):
escapedval = val.replace('\\', '\\\\')
if quote:
escapedval = escapedval.replace("'", "''")
escapedval = unicode_controlchars_re.sub(_show_control_chars, escapedval)
escapedval = UNICODE_CONTROLCHARS_RE.sub(_show_control_chars, escapedval)
bval = escapedval
if quote:
bval = "'{}'".format(bval)

View File

@ -19,7 +19,7 @@ import os
from .basecase import BaseTestCase
from .cassconnect import (get_cassandra_connection, create_keyspace, remove_db, testrun_cqlsh)
from cqlshlib.formatting import unicode_controlchars_re
from cqlshlib.formatting import UNICODE_CONTROLCHARS_RE
class TestCqlshUnicode(BaseTestCase):
@ -77,4 +77,4 @@ class TestCqlshUnicode(BaseTestCase):
self.assertIn(v2, output)
def test_unicode_esc(self): # CASSANDRA-17617
self.assertFalse(unicode_controlchars_re.match("01"))
self.assertFalse(UNICODE_CONTROLCHARS_RE.match("01"))

View File

@ -22,7 +22,6 @@ import os
import errno
import stat
from datetime import timedelta, tzinfo
from io import StringIO
try:
@ -31,21 +30,6 @@ try:
except ImportError:
HAS_LINE_PROFILER = False
ZERO = timedelta(0)
class UTC(tzinfo):
"""UTC"""
def utcoffset(self, dt):
return ZERO
def tzname(self, dt):
return "UTC"
def dst(self, dt):
return ZERO
def split_list(items, pred):
"""
@ -108,12 +92,6 @@ def identity(x):
return x
def trim_if_present(s, prefix):
if s.startswith(prefix):
return s[len(prefix):]
return s
def is_file_secure(filename):
try:
st = os.stat(filename)