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)

参考:

chttps://cloud.tencent.com/developer/ask/sof/107880701https://cloud.tencent.com/developer/ask/sof/107880701https://cloud.tencent.com/developer/ask/sof/107880701

Logo

码道开发者社区,聚焦华为云码道 CodeArts 代码智能体,沉淀 Agent、Skill、鸿蒙开发实战内容,供开发者查阅资料、交流技术、分享工程实践

更多推荐