使用Pyspark的HIVE JDBC连接将列名作为行值返回(不用重新打包)
·
from pyspark.sql import SparkSession
def fix_hive():
# 2. 获取类(不要 from)
ClassPool = jvm.javassist.ClassPool
CtClass = jvm.javassist.CtClass
CtConstructor = jvm.javassist.CtConstructor
CtField = jvm.javassist.CtField
CtNewMethod = jvm.javassist.CtNewMethod
AccessFlag = jvm.javassist.bytecode.AccessFlag
# 获取类池
pool = ClassPool.getDefault()
# 1. 创建新类:com.example.HiveDialect
cc = pool.makeClass("org.apache.spark.sql.jdbc.HiveDialect")
# ✅ 添加 public 无参构造函数
constructor = CtConstructor(None, cc) # 参数:参数类型列表(None 表示无参)
constructor.setModifiers(AccessFlag.PUBLIC)
# 2. 添加实现接口
jdbc_dialect = pool.get("org.apache.spark.sql.jdbc.JdbcDialect")
cc.setSuperclass(jdbc_dialect) # 如果是抽象类,否则用 addInterface
# 3. 添加单例实例(public static final HiveDialect INSTANCE)
instance_aa = CtField.make("public static final org.apache.spark.sql.jdbc.JdbcDialect INSTANCE;", cc)
cc.addField(instance_aa)
# instance_field = cc.addField(
# CtField(jpype.JClass("org.apache.spark.sql.jdbc.HiveDialect"), "INSTANCE", cc)
# )
# instance_field.setModifiers(AccessFlag.PUBLIC | AccessFlag.STATIC | AccessFlag.FINAL)
# cc.addField(instance_field)
# 4. 添加 canHandle 方法
can_handle_code = """
public boolean canHandle(String url) {
return url != null && url.startsWith("jdbc:hive2");
}
"""
can_handle_method = CtNewMethod.make(can_handle_code, cc)
cc.addMethod(can_handle_method)
# 5. 添加 quoteIdentifier 方法
quote_code = """
public String quoteIdentifier(String colName) {
if (colName == null || colName.isEmpty()) return colName;
String[] parts = colName.split("\\\\.");
StringBuilder sb = new StringBuilder();
for (int i = 0; i < parts.length; i++) {
sb.append("`").append(parts[i]).append("`");
if (i < parts.length - 1) sb.append(".");
}
return sb.toString();
}
"""
quote_method = CtNewMethod.make(quote_code, cc)
cc.addMethod(quote_method)
cc.toClass()
# 1. 创建 Spark 会话
spark = SparkSession.builder \
.appName("Stock Trend Prediction") \
.master("local[2]") \
.config("spark.jars",
"javassist-3.29.2-GA.jar,hive-jdbc-4.0.0-standalone.jar") \
.getOrCreate()
spark.sparkContext.setLogLevel("ERROR")
# 2. 获取 JVM 对象
jvm = spark.sparkContext._jvm
# 3.修复hiveDialect
fix_hive()
JdbcDialects = jvm.org.apache.spark.sql.jdbc.JdbcDialects
HiveDialect = jvm.org.apache.spark.sql.jdbc.HiveDialect
cc = HiveDialect()
JdbcDialects.registerDialect(cc)
query = """
(select
*
from stock_test) dummy_alias
"""
df = spark.read.format("jdbc") \
.option('url', 'jdbc:hive2://localhost:10000/default?hive.resultset.use.unique.column.names=false') \
.option("driver", "org.apache.hive.jdbc.HiveDriver") \
.option("user", "xxx") \
.option("dbtable", query) \
.option("sessionInitStatement", "SET hive.resultset.use.unique.column.names=false") \
.load()
print(df.count())
df.printSchema()
df.show(truncate=False)
参考:
更多推荐


所有评论(0)