· 9 years ago · Nov 02, 2016, 08:46 AM
1#include <atomic>
2#include <functional>
3#include <iostream>
4#include <map>
5#include <memory>
6#include <mutex>
7#include <thread>
8
9#include <boost/optional.hpp>
10
11#include <sqlite3.h>
12
13// Interfaces
14
15using ScopedDbConnection = std::unique_ptr<sqlite3, std::function<void(sqlite3*)>>;
16
17class DbConnectionManager {
18public:
19 virtual ScopedDbConnection getConnection() = 0;
20};
21
22class AbortTransaction {};
23class TransactionAborted {};
24
25class TransactionManager {
26public:
27 virtual void performInTransaction(const std::function<void()>& f) = 0;
28};
29
30struct Entity;
31
32class EntityRepoException {};
33
34class EntityRepo {
35public:
36 virtual void save(const Entity& entity) = 0;
37 virtual boost::optional<Entity> findById(uint64_t id) = 0;
38};
39
40// Implementations
41
42class ConcreteConnectionManager : public DbConnectionManager {
43public:
44 ConcreteConnectionManager(const std::string& dbFile);
45
46 ScopedDbConnection getConnection() override;
47private:
48 std::string _dbFile;
49};
50
51ConcreteConnectionManager::ConcreteConnectionManager(const std::string& dbFile) :
52 _dbFile(dbFile)
53{}
54
55ScopedDbConnection ConcreteConnectionManager::getConnection()
56{
57 sqlite3* connection;
58 int openResult = sqlite3_open_v2(
59 _dbFile.c_str(),
60 &connection,
61 SQLITE_OPEN_READWRITE | SQLITE_OPEN_CREATE | SQLITE_OPEN_NOMUTEX,
62 NULL
63 );
64 if (openResult != SQLITE_OK) {
65 std::cerr << "[-] Could not open database: " << openResult << std::endl;
66 sqlite3_close_v2(connection);
67 throw std::runtime_error("Database could not be opened");
68 }
69 return ScopedDbConnection(connection, [](sqlite3* c) { sqlite3_close_v2(c); });
70}
71
72
73
74class ConcreteTransactionManager : public DbConnectionManager, public TransactionManager {
75public:
76 ConcreteTransactionManager(ConcreteConnectionManager& connectionManager);
77
78 ScopedDbConnection getConnection() override;
79
80 void performInTransaction(const std::function<void()>& f) override;
81private:
82 struct TransactionInfo {
83 ScopedDbConnection dbConnection;
84 uint32_t count;
85 };
86
87 TransactionInfo& setupTransaction(const std::thread::id& threadId);
88 void abortTransaction(TransactionInfo& transactionInfo);
89 void commitTransaction(TransactionInfo& transactionInfo);
90
91 ConcreteConnectionManager& _connectionManager;
92 std::map<std::thread::id, TransactionInfo> _currentTransactions;
93 std::mutex _currentTransactionsMutex;
94};
95
96ConcreteTransactionManager::ConcreteTransactionManager(
97 ConcreteConnectionManager& connectionManager) :
98 _connectionManager(connectionManager)
99{}
100
101ScopedDbConnection ConcreteTransactionManager::getConnection()
102{
103 auto threadId = std::this_thread::get_id();
104 std::lock_guard<std::mutex> _(_currentTransactionsMutex);
105 auto it = _currentTransactions.find(threadId);
106 if (it == _currentTransactions.end()) {
107 return _connectionManager.getConnection();
108 }
109 return ScopedDbConnection(it->second.dbConnection.get(), [](sqlite3* c) {});
110}
111
112void ConcreteTransactionManager::abortTransaction(TransactionInfo& transactionInfo)
113{
114 int abortResult = sqlite3_exec(
115 transactionInfo.dbConnection.get(),
116 "ROLLBACK TRANSACTION",
117 nullptr,
118 nullptr,
119 nullptr
120 );
121
122 // If a transaction cannot be aborted, something bad happened.
123 assert(abortResult == SQLITE_OK);
124
125 auto threadId = std::this_thread::get_id();
126 {
127 std::lock_guard<std::mutex> _(_currentTransactionsMutex);
128 _currentTransactions.erase(threadId);
129 }
130}
131
132void ConcreteTransactionManager::commitTransaction(TransactionInfo& transactionInfo)
133{
134 int commitResult = sqlite3_exec(
135 transactionInfo.dbConnection.get(),
136 "COMMIT TRANSACTION",
137 nullptr,
138 nullptr,
139 nullptr
140 );
141
142 auto threadId = std::this_thread::get_id();
143 if (commitResult == SQLITE_OK) {
144 std::lock_guard<std::mutex> _(_currentTransactionsMutex);
145 _currentTransactions.erase(threadId);
146 return;
147 }
148
149 abortTransaction(transactionInfo);
150 throw TransactionAborted();
151}
152
153ConcreteTransactionManager::TransactionInfo&
154ConcreteTransactionManager::setupTransaction(const std::thread::id& threadId)
155{
156 std::lock_guard<std::mutex> _(_currentTransactionsMutex);
157 if (_currentTransactions.count(threadId) == 0) {
158 TransactionInfo transactionInfo{_connectionManager.getConnection(), 1};
159 int transactionStartResult = sqlite3_exec(
160 transactionInfo.dbConnection.get(),
161 "BEGIN TRANSACTION",
162 nullptr,
163 nullptr,
164 nullptr
165 );
166 if (transactionStartResult != SQLITE_OK) {
167 throw std::runtime_error("Could not start transaction");
168 }
169 auto it = _currentTransactions.emplace(threadId, std::move(transactionInfo)).first;
170 return it->second;
171 }
172 else {
173 TransactionInfo& transactionInfo = _currentTransactions[threadId];
174 ++transactionInfo.count;
175 return transactionInfo;
176 }
177}
178
179void ConcreteTransactionManager::performInTransaction(const std::function<void()>& f)
180{
181 auto threadId = std::this_thread::get_id();
182
183 TransactionInfo& transactionInfo = setupTransaction(threadId);
184
185 try {
186 f();
187 }
188 catch (AbortTransaction&) {
189 if (--transactionInfo.count > 0) {
190 throw;
191 }
192 abortTransaction(transactionInfo);
193 return;
194 }
195 catch (...) {
196 if (--transactionInfo.count > 0) {
197 throw;
198 }
199 abortTransaction(transactionInfo);
200 throw;
201 }
202
203 if (--transactionInfo.count > 0) {
204 return;
205 }
206 commitTransaction(transactionInfo);
207}
208
209struct Entity {
210 int64_t id;
211 int64_t balance;
212};
213
214class ConcreteEntityRepo : public EntityRepo {
215public:
216 ConcreteEntityRepo(DbConnectionManager& connectionManager);
217
218 void save(const Entity& entity) override;
219 boost::optional<Entity> findById(uint64_t id) override;
220private:
221 DbConnectionManager& _connectionManager;
222};
223
224ConcreteEntityRepo::ConcreteEntityRepo(DbConnectionManager& connectionManager) :
225 _connectionManager(connectionManager)
226{}
227
228void ConcreteEntityRepo::save(const Entity& entity)
229{
230 const char* sql = "INSERT OR REPLACE INTO entities (id, balance) VALUES (?, ?)";
231
232 ScopedDbConnection dbConnection = _connectionManager.getConnection();
233
234 sqlite3_stmt* rawStmt;
235 int stmtCreationResult = sqlite3_prepare_v2(
236 dbConnection.get(),
237 sql,
238 -1,
239 &rawStmt,
240 nullptr
241 );
242 if (stmtCreationResult != SQLITE_OK) {
243 throw EntityRepoException();
244 }
245 auto stmtDeleter = [](sqlite3_stmt* p) { sqlite3_finalize(p); };
246 std::unique_ptr<sqlite3_stmt, decltype(stmtDeleter)> stmt(rawStmt, stmtDeleter);
247
248 if (sqlite3_bind_int64(stmt.get(), 1, entity.id) != SQLITE_OK) {
249 throw EntityRepoException();
250 }
251 if (sqlite3_bind_int64(stmt.get(), 2, entity.balance) != SQLITE_OK) {
252 throw EntityRepoException();
253 }
254
255 if (sqlite3_step(stmt.get()) != SQLITE_DONE) {
256 throw EntityRepoException();
257 }
258}
259
260boost::optional<Entity> ConcreteEntityRepo::findById(uint64_t id)
261{
262 const char* sql = "SELECT id, balance FROM entities WHERE id = ?";
263
264 ScopedDbConnection dbConnection = _connectionManager.getConnection();
265
266 sqlite3_stmt* rawStmt;
267 int stmtCreationResult = sqlite3_prepare_v2(
268 dbConnection.get(),
269 sql,
270 -1,
271 &rawStmt,
272 nullptr
273 );
274 if (stmtCreationResult != SQLITE_OK) {
275 throw EntityRepoException();
276 }
277 auto stmtDeleter = [](sqlite3_stmt* p) { sqlite3_finalize(p); };
278 std::unique_ptr<sqlite3_stmt, decltype(stmtDeleter)> stmt(rawStmt, stmtDeleter);
279
280 if (sqlite3_bind_int64(stmt.get(), 1, id) != SQLITE_OK) {
281 throw EntityRepoException();
282 }
283
284 int stepResult = sqlite3_step(stmt.get());
285 if (stepResult == SQLITE_ROW) {
286 Entity result{sqlite3_column_int64(stmt.get(), 0), sqlite3_column_int64(stmt.get(), 1)};
287 return result;
288 }
289 else if (stepResult == SQLITE_DONE) {
290 return boost::none;
291 }
292 else {
293 throw EntityRepoException();
294 }
295}
296
297void noTransaction(EntityRepo& repo)
298{
299 Entity e{1, 1000};
300 repo.save(e);
301 std::cout << "[+] No transaction: finished" << std::endl;
302}
303
304void inTransactionNoOneElse(TransactionManager& transactionManager, EntityRepo& entityRepo)
305{
306 transactionManager.performInTransaction([&]() {
307 Entity e{2, 1000};
308 entityRepo.save(e);
309 Entity e2 = *entityRepo.findById(2);
310 e2.balance += 1000;
311 entityRepo.save(e2);
312 });
313
314 std::cout << "[+] In transaction with no one else" << std::endl;
315}
316
317void transactionAborted(TransactionManager& transactionManager, EntityRepo& entityRepo)
318{
319 transactionManager.performInTransaction([&]() {
320 Entity e{3, 1000};
321 entityRepo.save(e);
322 Entity e2 = *entityRepo.findById(3);
323 e2.balance += 1000;
324 entityRepo.save(e2);
325 throw AbortTransaction();
326 });
327
328 std::cout << "[+] Abort transaction example: Finished" << std::endl;
329}
330
331void simultaneousTransactions(TransactionManager& transactionManager, EntityRepo& entityRepo)
332{
333 std::atomic<bool> saveExecuted{false};
334 std::atomic<bool> findExecuted{false};
335
336 std::thread t1([&]() {
337 transactionManager.performInTransaction([&]() {
338 while (!saveExecuted) {}
339 boost::optional<Entity> e = entityRepo.findById(4);
340 findExecuted = true;
341 if (e) {
342 std::cout << "[-] Transaction behavior violated!" << std::endl;
343 }
344 else {
345 std::cout << "[+] Simultaneous non-interfering transactions: success" << std::endl;
346 }
347 });
348 });
349
350 std::thread t2([&]() {
351 transactionManager.performInTransaction([&]() {
352 Entity e{4, 1000};
353 entityRepo.save(e);
354 saveExecuted = true;
355 while (!findExecuted) {}
356 });
357 });
358
359 t1.join();
360 t2.join();
361
362 std::cout << "[+] Simultaneous transaction example: Finished" << std::endl;
363}
364
365void simultaneousConflictingTransactions(
366 TransactionManager& transactionManager,
367 EntityRepo& entityRepo)
368{
369 std::atomic<bool> saveExecuted{false};
370 std::atomic<bool> findExecuted{false};
371
372 std::thread t1([&]() {
373 transactionManager.performInTransaction([&]() {
374 while (!saveExecuted) {}
375 boost::optional<Entity> e = entityRepo.findById(5);
376 if (e) {
377 std::cout << "[-] Transaction behavior violated!" << std::endl;
378 }
379 else {
380 try {
381 entityRepo.save(Entity{5, 2000});
382 }
383 catch (EntityRepoException&) {
384 findExecuted = true;
385 std::cout << "[+] Conflicting operations: correct" << std::endl;
386 throw AbortTransaction();
387 }
388 }
389
390 });
391 });
392
393 std::thread t2([&]() {
394 transactionManager.performInTransaction([&]() {
395 Entity e{5, 1000};
396 entityRepo.save(e);
397 saveExecuted = true;
398 while (!findExecuted) {}
399 });
400 });
401
402 t1.join();
403 t2.join();
404
405 std::cout << "[+] Simultaneous transaction example: Finished" << std::endl;
406}
407
408void simultaneousConflictingTransactions2(
409 TransactionManager& transactionManager,
410 EntityRepo& entityRepo)
411{
412 std::atomic<bool> saveExecuted{false};
413 std::atomic<bool> findExecuted{false};
414 std::atomic<bool> transactionCommitted{false};
415
416 std::thread t1([&]() {
417 transactionManager.performInTransaction([&]() {
418 while (!saveExecuted) {}
419 boost::optional<Entity> e = entityRepo.findById(6);
420 findExecuted = true;
421 if (e) {
422 std::cout << "[-] Transaction behavior violated!" << std::endl;
423 throw AbortTransaction();
424 }
425 while (!transactionCommitted) {}
426 try {
427 entityRepo.save(Entity{6, 2000});
428 }
429 catch (EntityRepoException&) {
430 std::cout << "[+] Conflicting operations: correct" << std::endl;
431 throw AbortTransaction();
432 }
433 });
434 });
435
436 std::thread t2([&]() {
437 transactionManager.performInTransaction([&]() {
438 Entity e{6, 1000};
439 entityRepo.save(e);
440 saveExecuted = true;
441 while (!findExecuted) {}
442 });
443
444 transactionCommitted = true;
445 });
446
447 t1.join();
448 t2.join();
449
450 std::cout << "[+] Simultaneous transaction example: Finished" << std::endl;
451}
452
453void nestedTransaction(TransactionManager& transactionManager, EntityRepo& entityRepo)
454{
455 transactionManager.performInTransaction([&]() {
456 Entity e{7, 1000};
457 entityRepo.save(e);
458 transactionManager.performInTransaction([&]() {
459 boost::optional<Entity> e2 = entityRepo.findById(7);
460 assert(e2);
461 e2->balance += 1000;
462 entityRepo.save(*e2);
463 });
464 });
465
466 std::cout << "[+] Nested transaction: Finished" << std::endl;
467}
468
469void nestedTransactionOuterAbort(TransactionManager& transactionManager, EntityRepo& entityRepo)
470{
471 transactionManager.performInTransaction([&]() {
472 Entity e{8, 1000};
473 entityRepo.save(e);
474 transactionManager.performInTransaction([&]() {
475 boost::optional<Entity> e2 = entityRepo.findById(8);
476 assert(e2);
477 e2->balance += 1000;
478 entityRepo.save(*e2);
479 });
480 throw AbortTransaction();
481 });
482
483 std::cout << "[+] Nested transaction, outer aborts: Finished" << std::endl;
484}
485
486void nestedTransactionInnerAbort(TransactionManager& transactionManager, EntityRepo& entityRepo)
487{
488 transactionManager.performInTransaction([&]() {
489 Entity e{9, 1000};
490 entityRepo.save(e);
491 transactionManager.performInTransaction([&]() {
492 boost::optional<Entity> e2 = entityRepo.findById(9);
493 assert(e2);
494 e2->balance += 1000;
495 entityRepo.save(*e2);
496 throw AbortTransaction();
497 });
498 });
499
500 std::cout << "[+] Nested transaction, inner aborts: Finished" << std::endl;
501}
502
503
504int main(int argc, char** argv)
505{
506 const std::string dbFile = "poc.db";
507 ConcreteConnectionManager connectionManager(dbFile);
508 ConcreteTransactionManager transactionManager(connectionManager);
509 ConcreteEntityRepo entityRepo(transactionManager);
510
511 {
512 ScopedDbConnection conn = connectionManager.getConnection();
513
514 const char* dropTableSql = "DROP TABLE IF EXISTS entities";
515 sqlite3_exec(conn.get(), dropTableSql, nullptr, nullptr, nullptr);
516
517 const char* createTableSql = "CREATE TABLE IF NOT EXISTS entities "
518 "(id INT PRIMARY KEY, balance INT)";
519 sqlite3_exec(conn.get(), createTableSql, nullptr, nullptr, nullptr);
520
521 const char* walModePragma = "PRAGMA journal_mode=WAL";
522 sqlite3_exec(conn.get(), walModePragma, nullptr, nullptr, nullptr);
523
524 std::cout << "[+] DB Created" << std::endl;
525 }
526 // No transaction
527 noTransaction(entityRepo);
528
529 inTransactionNoOneElse(transactionManager, entityRepo);
530
531 transactionAborted(transactionManager, entityRepo);
532
533 simultaneousTransactions(transactionManager, entityRepo);
534
535 simultaneousConflictingTransactions(transactionManager, entityRepo);
536
537 simultaneousConflictingTransactions2(transactionManager, entityRepo);
538
539 nestedTransaction(transactionManager, entityRepo);
540
541 nestedTransactionOuterAbort(transactionManager, entityRepo);
542
543 nestedTransactionInnerAbort(transactionManager, entityRepo);
544}