From d8da4382473a6d1b86a2d934ef2632fd24512652 Mon Sep 17 00:00:00 2001 From: Jonathan Zernik Date: Fri, 31 Jul 2020 22:08:19 -0700 Subject: [PATCH] Fix setting up schema in createdb and postgres params (#189) --- createdb.sql | 5 +++++ squeakserver/server/main.py | 16 ++++++++-------- squeakserver/server/postgres_db.py | 25 +------------------------ 3 files changed, 14 insertions(+), 32 deletions(-) diff --git a/createdb.sql b/createdb.sql index 24618bd8..501441e7 100644 --- a/createdb.sql +++ b/createdb.sql @@ -1,3 +1,8 @@ CREATE DATABASE squeakserver; \connect squeakserver; + +CREATE SCHEMA IF NOT EXISTS mainnet; +CREATE SCHEMA IF NOT EXISTS testnet; +CREATE SCHEMA IF NOT EXISTS regtest; +CREATE SCHEMA IF NOT EXISTS simnet; diff --git a/squeakserver/server/main.py b/squeakserver/server/main.py index aafbdb13..5301f3f9 100644 --- a/squeakserver/server/main.py +++ b/squeakserver/server/main.py @@ -71,12 +71,13 @@ def load_admin_handler(lightning_client, squeak_node): return SqueakAdminServerHandler(lightning_client, squeak_node,) -def load_db_params(config): - return parse_db_params(config) - - -def load_postgres_db(config): +def load_db_params(config, schema): db_params = parse_db_params(config) + db_params['options'] = "--search_path={}".format(schema) + return db_params + + +def load_postgres_db(db_params): return PostgresDb(db_params) @@ -144,14 +145,13 @@ def run_server(config): SelectParams(network) # load the db params - db_params = load_db_params(config) + db_params = load_db_params(config, network) logger.info("db params: " + str(db_params)) # load postgres db - postgres_db = load_postgres_db(config) + postgres_db = load_postgres_db(db_params) logger.info("postgres_db: " + str(postgres_db)) postgres_db.get_version() - postgres_db.use_schema(network) postgres_db.init() # load the price diff --git a/squeakserver/server/postgres_db.py b/squeakserver/server/postgres_db.py index a9f95b30..62b9de0d 100644 --- a/squeakserver/server/postgres_db.py +++ b/squeakserver/server/postgres_db.py @@ -18,12 +18,8 @@ logger = logging.getLogger(__name__) class PostgresDb: def __init__(self, params): + logger.info("Starting connection pool with params: {}".format(params)) self.connection_pool = pool.ThreadedConnectionPool(5, 20, **params) - self.params = params - - @property - def user(self): - return self.params['user'] # Get Cursor @contextmanager @@ -46,25 +42,6 @@ class PostgresDb: db_version = curs.fetchone() logger.info(db_version) - def use_schema(self, schema_name): - """ Create the schema for the given name. """ - create_schema_sql = sql.SQL("CREATE SCHEMA IF NOT EXISTS {};").format( - sql.Identifier(schema_name) - ) - use_schema_sql = sql.SQL("SET search_path TO {};").format( - sql.Identifier(schema_name), - ) - use_schema_for_user_sql = sql.SQL("ALTER ROLE {} SET search_path TO {};").format( - sql.Identifier(self.user), - sql.Identifier(schema_name), - ) - with self.get_cursor() as curs: - # execute a statement - logger.info("Creating schema: {}".format(schema_name)) - curs.execute(create_schema_sql) - curs.execute(use_schema_sql) - curs.execute(use_schema_for_user_sql) - def init(self): """ Create the tables and indices in the database. """ with self.get_cursor() as curs: