PostgreSQLScriptRunner.java

/*
** Module   : PostgreSQLScriptRunner.java
** Abstract : Run DDL SQL script against the PostgreSQL database
**
** Copyright (c) 2022-2023, Golden Code Development Corporation.
**
** -#- -I- --Date-- ---------------------------------------Description---------------------------------------
** 001 IAS 20220816 Created initial version.
*  002 IAS 20230908 Add support for more SQL script types.
*/

/*
** This program is free software: you can redistribute it and/or modify
** it under the terms of the GNU Affero General Public License as
** published by the Free Software Foundation, either version 3 of the
** License, or (at your option) any later version.
**
** This program is distributed in the hope that it will be useful,
** but WITHOUT ANY WARRANTY; without even the implied warranty of
** MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE.  See the
** GNU Affero General Public License for more details.
**
** You may find a copy of the GNU Affero GPL version 3 at the following
** location: https://www.gnu.org/licenses/agpl-3.0.en.html
** 
** Additional terms under GNU Affero GPL version 3 section 7:
** 
**   Under Section 7 of the GNU Affero GPL version 3, the following additional
**   terms apply to the works covered under the License.  These additional terms
**   are non-permissive additional terms allowed under Section 7 of the GNU
**   Affero GPL version 3 and may not be removed by you.
** 
**   0. Attribution Requirement.
** 
**     You must preserve all legal notices or author attributions in the covered
**     work or Appropriate Legal Notices displayed by works containing the covered
**     work.  You may not remove from the covered work any author or developer
**     credit already included within the covered work.
** 
**   1. No License To Use Trademarks.
** 
**     This license does not grant any license or rights to use the trademarks
**     Golden Code, FWD, any Golden Code or FWD logo, or any other trademarks
**     of Golden Code Development Corporation. You are not authorized to use the
**     name Golden Code, FWD, or the names of any author or contributor, for
**     publicity purposes without written authorization.
** 
**   2. No Misrepresentation of Affiliation.
** 
**     You may not represent yourself as Golden Code Development Corporation or FWD.
** 
**     You may not represent yourself for publicity purposes as associated with
**     Golden Code Development Corporation, FWD, or any author or contributor to
**     the covered work, without written authorization.
** 
**   3. No Misrepresentation of Source or Origin.
** 
**     You may not represent the covered work as solely your work.  All modified
**     versions of the covered work must be marked in a reasonable way to make it
**     clear that the modified work is not originating from Golden Code Development
**     Corporation or FWD.  All modified versions must contain the notices of
**     attribution required in this license.
*/

package com.goldencode.p2j.persist.deploy;

import java.io.*;
import java.sql.*;
import java.util.*;
import java.util.logging.*;
import java.util.regex.*;
import java.util.stream.*;

import com.goldencode.p2j.cfg.*;
import com.goldencode.p2j.persist.*;
import com.goldencode.p2j.persist.dialect.*;
import com.goldencode.p2j.util.*;

/** 
 *  Run SQL script against the PostgreSQL database
 *  
 */
public class PostgreSQLScriptRunner 
extends ScriptRunner
{
   /** pattern for the PostgreSQL 'SELECT version()' result parsing */
   private static final Pattern PG_VER = Pattern.compile("^PostgreSQL *([0-9]*)\\.([0-9]*).*");
   

   /**
    *  PostgreSQL query for counting the number of UDFs with a spicefic name.
    */
   private static final List<String> COUNT = Collections.unmodifiableList(
            Arrays.asList(
               "select count(*)",
               "from pg_proc p",
               "left join pg_namespace n on p.pronamespace = n.oid",
               "left join pg_language l on p.prolang = l.oid",
               "left join pg_type t on t.oid = p.prorettype",
               "where n.nspname = ?",
               "  and l.lanname in ('sql', 'plpgsql')",
               "  and t.typname != 'trigger'",
               "  and p.proname = ?"
            )
         );

   /** SCALE UDF (built-in for PostgreSQL 10+) */
   private static final List<String> SCALE = Collections.unmodifiableList(
      Arrays.asList(
         "CREATE OR REPLACE FUNCTION scale(v numeric)",
         "RETURNS integer",
         "LANGUAGE plpgsql",
         "IMMUTABLE STRICT LEAKPROOF",
         "AS $function$",
         "   declare s text;",
         "   declare p integer;",
         "   begin",
         "      s := v::text;",
         "      p := strpos(s, '.');",
         "      if p = 0 then",
         "         return 0;",
         "      else",
         "         return length(left(s, length(s) - p));",
         "      end if;",
         "   end",
         "$function$"
      )
   );

   /**
    * Constructor.
    * 
    * @param dialect
    *        the target database dialect 
    */
   public PostgreSQLScriptRunner(Dialect dialect)
   {
      super(dialect);
   }

   
   /**
    * Apply UDF SQL scripts.
    * 
    * @param conn
    *        Database connection.
    *        
    * @throws SQLException
    *         on SQL error.
    * @throws IOException
    *         on script reading error.
    */
   @Override
   protected void applyUdfScripts(Connection conn) 
   throws SQLException, IOException
   {
      ScriptSplitter splitter = dialect.scriptSplitter();

      
       // Creating 'scale' UDF if PostgreSQL version is less than 10.00
      if (getPgVersion(conn) < 1000)
      {
         createScaleUDF(conn);
      }
      
      applyScripts(conn, splitter);

      // Setting the database 'search_path'


      if (scriptType == ScriptType.UDF_INSTALL_SET_SEARCH_PATH)
      {
         String dbname = null;
         try(Statement stmt = conn.createStatement();
             ResultSet rs = stmt.executeQuery("SELECT current_database()"))
         {
            if( rs.next())
            {
               dbname = rs.getString(1);
            }
         }
         if (dbname != null)
         {
            LOG.log(Level.INFO, String.format("Setting search_path for [%s]", dbname));
            execute(conn, "ALTER DATABASE " + dbname + " SET search_path TO public,udf");
         }
         else
         {
            LOG.log(Level.WARNING, "Failed to retrieve the database name");
         }
      }
      
   }
   
   /**
    * Check for missing UDFs and create them.
    * 
    * @param dbname
    *        Database name.
    * @param url
    *         Database JDBC URL.
    * @param adm
    *        Admin login.
    * @param pwd
    *        Admin password. 
    *        
    * @throws SQLException
    *         on SQL error.
    * @throws PersistenceException
    *         on script reading error.
    */
   @Override
   protected void createMissingUdfs(String dbname, String url, String adm, String pwd) 
   throws SQLException, 
          PersistenceException
   {
      List<String> scriptNames = new ArrayList<>();
      try(Connection conn = DriverManager.getConnection(url, adm, pwd))
      {
         if (countUDFs(conn, "udf", "getfwdversion") == 0)
         {
            LOG.info(String.format("%s UDF not found in %s; %s SQL script will be applied", 
                     "'udf.getfwdversion'", dbname, "'udfs.sql'"));
            scriptNames.add("udfs.sql");
         }
         if (countUDFs(conn, "public", "words") < 2)
         {
            LOG.info(String.format("Less than 2 %s UDFs found in %s; %s SQL script will be applied", 
                     "'public.words'", dbname, "'words-udfs-sql.sql'"));
            scriptNames.add("words-udfs-sql.sql");
         }
         scripts = addScriptPath(scriptNames);
         try
         {
            applyUdfScripts(conn);
         }
         catch (IOException e)
         {
            throw new PersistenceException(e);
         }
      }
   }
   
   /**
    * Retrieve PostgreSQL server version.
    *
    * @param   conn
    *          The database connection.
    *
    * @return  PostgreSQL server version as 100+major + minor.
    *
    * @throws  SQLException
    *          On error.
    */
   private static int getPgVersion(Connection conn)
   throws SQLException
   {
      String version = "?";
      try (Statement stmt = conn.createStatement())
      {
         ResultSet rs = stmt.executeQuery("SELECT version()");
         if( rs.next())
         {
            version = rs.getString(1);
         }
      }
      Matcher matcher = PG_VER.matcher(version);
      if (!matcher.matches())
      {
         throw new IllegalStateException("Unexpected 'SELECT version()' result: [" + version + "]");
      }
      return Integer.parseInt(matcher.group(1)) * 100 + Integer.parseInt(matcher.group(2)); 
   }

   /**
    * Count number of UDFs with a given name.
    *
    * @param   conn
    *          The database connection.
    * @param   schema
    *          schema name.
    * @param   name
    *          UDF name.
    * 
    * @return  number of UDFs with a given name.  
    *        
    * @throws  SQLException
    *          On error.
    */
   private int countUDFs(Connection conn, String schema, String name)
   throws SQLException
   {
      String eoln = EnvironmentOps.OS_WIN.equalsIgnoreCase(Configuration.getParameter("opsys"))
            ? "\r\n" /* express to WINDOWS */ : "\n" /* default to Linux */;  
      String sql = COUNT.stream().collect(Collectors.joining(eoln));
      try (PreparedStatement pstmt = conn.prepareStatement(sql))
      {
         int n = 0;
         pstmt.setString(1, schema);
         pstmt.setString(2, name);
         ResultSet rs = pstmt.executeQuery();
         if (rs.next())
         {
            n = rs.getInt(1);
         }
         return n;
      }
   }

   /**
    * Create SCALE UDF.
    *
    * @param   conn
    *          The database connection.
    *
    * @throws  SQLException
    *          On error.
    */
   private void createScaleUDF(Connection conn)
   throws SQLException
   {
      String eoln = EnvironmentOps.OS_WIN.equalsIgnoreCase(Configuration.getParameter("opsys"))
            ? "\r\n" /* express to WINDOWS */ : "\n" /* default to Linux */;  
      String sql = SCALE.stream().collect(Collectors.joining(eoln));
      execute(conn, sql);
   }
}