Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
13 changes: 8 additions & 5 deletions lib/core/replication.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,9 @@
from lib.core.settings import UNICODE_ENCODING
from lib.utils.safe2bin import safechardecode

def _quoteIdentifier(name):
return '"%s"' % name.replace('"', '""')

class Replication(object):
"""
This class holds all methods/classes used for database
Expand Down Expand Up @@ -62,11 +65,11 @@ def __init__(self, parent, name, columns=None, create=True, typeless=False):
self.columns = columns
if create:
try:
self.execute('DROP TABLE IF EXISTS "%s"' % self.name)
self.execute('DROP TABLE IF EXISTS %s' % _quoteIdentifier(self.name))
if not typeless:
self.execute('CREATE TABLE "%s" (%s)' % (self.name, ','.join('"%s" %s' % (unsafeSQLIdentificatorNaming(colname), coltype) for colname, coltype in self.columns)))
self.execute('CREATE TABLE %s (%s)' % (_quoteIdentifier(self.name), ','.join('%s %s' % (_quoteIdentifier(unsafeSQLIdentificatorNaming(colname)), coltype) for colname, coltype in self.columns)))
else:
self.execute('CREATE TABLE "%s" (%s)' % (self.name, ','.join('"%s"' % unsafeSQLIdentificatorNaming(colname) for colname in self.columns)))
self.execute('CREATE TABLE %s (%s)' % (_quoteIdentifier(self.name), ','.join(_quoteIdentifier(unsafeSQLIdentificatorNaming(colname)) for colname in self.columns)))
except Exception as ex:
errMsg = "problem occurred ('%s') while initializing the sqlite database " % getSafeExString(ex, UNICODE_ENCODING)
errMsg += "located at '%s'" % self.parent.dbpath
Expand All @@ -78,7 +81,7 @@ def insert(self, values):
"""

if len(values) == len(self.columns):
self.execute('INSERT INTO "%s" VALUES (%s)' % (self.name, ','.join(['?'] * len(values))), safechardecode(values))
self.execute('INSERT INTO %s VALUES (%s)' % (_quoteIdentifier(self.name), ','.join(['?'] * len(values))), safechardecode(values))
else:
errMsg = "wrong number of columns used in replicating insert"
raise SqlmapValueException(errMsg)
Expand Down Expand Up @@ -109,7 +112,7 @@ def select(self, condition=None):
"""
This function is used for selecting row(s) from current table.
"""
query = 'SELECT * FROM "%s"' % self.name
query = 'SELECT * FROM %s' % _quoteIdentifier(self.name)
if condition:
query += ' WHERE %s' % condition

Expand Down
18 changes: 18 additions & 0 deletions tests/test_dump_format.py
Original file line number Diff line number Diff line change
Expand Up @@ -328,6 +328,24 @@ def test_html_escapes_markup(self):


class TestSqliteDump(_FileDumpCase):
def test_double_quotes_in_table_and_column_names(self):
tv = _PlainOrderedDict([
("__infos__", {"count": 1, "db": "testdb", "table": '`sales"archive`'}),
('`unit"price`', {"length": 4, "values": ["3.50"]}),
])
conf.dumpFormat = DUMP_FORMAT.SQLITE
self.d.dbTableValues(tv)

import sqlite3
conn = sqlite3.connect(os.path.join(self.tmp, "testdb.sqlite3"))
try:
rows = conn.execute('SELECT "unit""price" FROM "sales""archive"').fetchall()
self.assertEqual(rows, [("3.50",)])
columns = conn.execute('PRAGMA table_info("sales""archive")').fetchall()
self.assertEqual(columns[0][1], 'unit"price')
finally:
conn.close()

def test_rows_and_inferred_types(self):
tv = _PlainOrderedDict([
("__infos__", {"count": 2, "db": "testdb", "table": "people"}),
Expand Down
54 changes: 54 additions & 0 deletions tests/test_replication.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,8 @@
from _testutils import bootstrap
bootstrap()

from lib.core.common import Backend
from lib.core.data import conf, kb
from lib.core.replication import Replication
from lib.core.exception import SqlmapConnectionException
from lib.core.exception import SqlmapValueException
Expand Down Expand Up @@ -96,6 +98,58 @@ def test_wrong_column_count_raises(self):
self.assertEqual(t.select(), [(1, "x")])


class TestQuotedIdentifiers(_ReplCase):
def setUp(self):
super(TestQuotedIdentifiers, self).setUp()
self._savedConf = dict((k, conf.get(k)) for k in ("forceDbms", "dbms"))
self._savedKb = dict((k, kb.get(k)) for k in ("forcedDbms", "dbms"))
conf.forceDbms = conf.dbms = None
kb.dbms = None
Backend.forceDbms("MySQL")

def tearDown(self):
for k, v in self._savedConf.items():
conf[k] = v
for k, v in self._savedKb.items():
kb[k] = v
super(TestQuotedIdentifiers, self).tearDown()

def test_table_name_with_double_quote(self):
t = self.rep.createTable('sales"archive', [("id", self.rep.INTEGER)])
self.assertEqual(t.name, 'sales"archive')
t.insert([1])
self.assertEqual(t.select(), [(1,)])
self.assertEqual(self._readback('SELECT id FROM "sales""archive"'), [(1,)])
replacement = self.rep.createTable('sales"archive', [("id", self.rep.INTEGER)])
self.assertEqual(replacement.select(), [])

def test_typed_column_name_with_double_quote(self):
t = self.rep.createTable("t", [('unit"price', self.rep.REAL)])
t.insert([3.5])
self.assertEqual(t.select(), [(3.5,)])
columns = self._readback("PRAGMA table_info(t)")
self.assertEqual(columns[0][1:3], ('unit"price', "REAL"))

def test_typeless_column_name_with_double_quote(self):
t = self.rep.createTable("t", ['unit"price'], typeless=True)
t.insert(["3.50"])
self.assertEqual(t.select(), [("3.50",)])
columns = self._readback("PRAGMA table_info(t)")
self.assertEqual(columns[0][1:3], ('unit"price', ""))

def test_insert_into_existing_table_with_double_quote(self):
self.rep.connection.execute('CREATE TABLE "sales""archive" (id INTEGER)')
t = Replication.Table(self.rep, 'sales"archive', [("id", self.rep.INTEGER)], create=False)
t.insert([7])
self.assertEqual(self._readback('SELECT id FROM "sales""archive"'), [(7,)])

def test_select_from_existing_table_with_double_quote(self):
self.rep.connection.execute('CREATE TABLE "sales""archive" (id INTEGER)')
self.rep.connection.execute('INSERT INTO "sales""archive" VALUES (7)')
t = Replication.Table(self.rep, 'sales"archive', [("id", self.rep.INTEGER)], create=False)
self.assertEqual(t.select("id = 7"), [(7,)])


class TestInitFailure(unittest.TestCase):
"""A failed open (e.g. unwritable path) must raise cleanly and the partially
constructed object must be safe to finalize (no AttributeError in __del__)."""
Expand Down