Spark3.x指北——3:Spark SQL
本章为Spark SQL详细内容~
Spark3.x指北全系列目录:
Spark基础概念请看:Spark3.x指北——1:Spark基础概念
SparkCore内容请看:Spark3.x指北——2:Spark Core
SparkStreaming内容请看:Spark3.x指北——4:SparkStreaming
目录
6.4 DataFrame & DataSet的使用(spark-shell)
(2)使用SQL语法操作DataFrame——通过spark.sql
(3)DSL语法——通过DataFrame直接调用DSL(类似方法调用)
a RDD To DataFrame——通过RDD.toDF(字段1, 字段2)
a DataFrame To DataSet——通过DF.as[T]
b DataSet To DataFrame——通过DS.toDF
6.4.3 RDD & DataFrame & DataSet
6.5.2 创建SparkSQL上下文环境对象(类似Spark Core中的SparkContext)
6.5.5 RDD – DataFrame – DataSet之间的相互转换
(2)UDF函数使用——通过spark.udf.register直接注册
(3)通过弱类型实现UDAF求age字段均值——继承UserDefinedAggregateFunction类
a 首先我们需要继承并重写UserDefinedAggregateFunctio类
⑧ evaluate——以最终的缓冲区结果进行聚合计算,得出最终结果
b 在sparkSQL中通过spark.udf注册并使用UDAF
c 另一种StructType构造(通过StructField :: Nil)实现UDAF的参考
a 同样的,我们需要继承Aggregator[IN, BUF, OUT]来实现强类型的UDAF
③ reduce——根据输入数据(查询到的新数据)更新缓冲区
④ merge——跨分区shuffle,合并不同分区的缓冲区
⑥ bufferEncoder & outputEncoder——指定缓冲区 & 输出数据编码格式
(1)spark.read.load(“文件”)——通用读取
(2)操作外部已部署好的Hive——通过spark-shell
① 将hive-site.xml拷贝到当前spark的conf/目录下
(3)操作外部已部署好的Hive——通过IDEA在代码中操作
b 将hive-site.xml拷贝到当前项目的resources文件夹下,作为连接hive的配置文件
③ 按找地区进行分组,对每个地区的点击次数降序排列,取出每个地区的点击次数前三
b 对分组sql进行修改,在按区域、商品分组时,通过UDAF统计城市备注信息
6 Spark SQL

6.1 Spark SQL概述
SparkSQL 的前身是Shark,给熟悉RDBMS但又不理解MapReduce的技术人员提供快 速上手的工具。 Hive 是早期唯一运行在Hadoop上的SQL-on-Hadoop工具。但是MapReduce计算过程 中大量的中间磁盘落地过程消耗了大量的I/O,降低的运行效率,为了提高SQL-on-Hadoop 的效率,大量的SQL-on-Hadoop工具开始产生,其中表现较为突出的是:
- Drill
- Impala
- Shark
其中Shark是伯克利实验室Spark生态环境的组件之一,是基于Hive所开发的工具,它修改了下图所示的右下角的内存管理、物理计划、执行三个模块,并使之能运行在Spark引擎 上。

Shark 的出现,使得SQL-on-Hadoop的性能比Hive有了10-100倍的提高。

但是,随着Spark的发展,对于野心勃勃的Spark团队来说,Shark对于Hive的太多依 赖(如采用Hive的语法解析器、查询优化器等等),制约了Spark的One Stack Rule Them All 的既定方针,制约了Spark各个组件的相互集成,所以提出了SparkSQL项目。SparkSQL 抛弃原有Shark的代码,汲取了Shark的一些优点,如内存列存储(In-Memory Columnar Storage)、Hive兼容性等,重新开发了SparkSQL代码;由于摆脱了对Hive的依赖性,SparkSQL 无论在数据兼容、性能优化、组件扩展方面都得到了极大的方便,真可谓“退一步,海阔天 空”。
- 数据兼容方面 SparkSQL不但兼容Hive,还可以从RDD、parquet文件、JSON文件中 获取数据,未来版本甚至支持获取RDBMS数据以及cassandra等NOSQL数据;
- 性能优化方面 除了采取In-Memory Columnar Storage、byte-code generation 等优化技术 外、将会引进Cost Model对查询进行动态评估、获取最佳物理计划等等;
- 组件扩展方面 无论是SQL的语法解析器、分析器还是优化器都可以重新定义,进行扩 展。

2014 年6月1日Shark项目和SparkSQL项目的主持人Reynold Xin宣布:停止对Shark的 开发,团队将所有资源放SparkSQL项目上,至此,Shark的发展画上了句话,但也因此发 展出两个支线:SparkSQL和Hive on Spark。

其中SparkSQL作为Spark生态的一员继续发展,而不再受限于Hive,只是兼容Hive;而 Hive on Spark 是一个 Hive 的发展计划,该计划将Spark作为Hive的底层引擎之一,也就是 说,Hive将不再受限于一个引擎,可以采用Map-Reduce、Tez、Spark等引擎。
对于开发人员来讲,SparkSQL可以简化RDD的开发,提高开发效率,且执行效率非常快,所以实际工作中,基本上采用的就是SparkSQL。Spark SQL为了简化RDD的开发, 提高开发效率,提供了2个编程抽象,类似Spark Core中的RDD
- DataFrame
- DataSet
6.2 Spark SQL特点
- 易整合 无缝的整合了 SQL 查询和 Spark 编程

- 统一的数据访问 使用相同的方式连接不同的数据源

- 兼容Hive 在已有的仓库上直接运行 SQL 或者 HiveQL

- 标准数据连接 通过 JDBC 或者 ODBC 来连接

6.3 DataFrame & DataSet概述
6.3.1 DataFrame
在Spark 中,DataFrame 是一种以RDD为基础的分布式数据集,类似于传统数据库中 的二维表格。DataFrame与RDD的主要区别在于,前者带有schema元信息,即DataFrame 所表示的二维表数据集的每一列都带有名称和类型。这使得Spark SQL得以洞察更多的结构 信息,从而对藏于DataFrame背后的数据源以及作用于DataFrame之上的变换进行了针对性的优化,最终达到大幅提升运行时效率的目标。反观RDD,由于无从得知所存数据元素的 具体内部结构,Spark Core只能在stage层面进行简单、通用的流水线优化。
简单来说,RDD只关心数据本身,而DataFrame同时关心数据本身以及数据结构,这使得SparkSQL可以根据结构进行执行优化。
同时,与Hive类似,DataFrame也支持嵌套数据类型(struct、array和map)。从 API 易用性的角度上看,DataFrame API提供的是一套高层的关系操作,比函数式的RDD API 要 更加友好,门槛更低。

上图直观地体现了DataFrame和RDD的区别。 左侧的RDD[Person]虽然以Person为类型参数,但Spark框架本身不了解Person类的内 部结构。而右侧的DataFrame却提供了详细的结构信息,使得 Spark SQL 可以清楚地知道 该数据集中包含哪些列,每列的名称和类型各是什么。 DataFrame 是为数据提供了Schema的视图。可以把它当做数据库中的一张表来对待。
DataFrame 也是懒执行的,但性能上比RDD要高,主要原因:优化的执行计划,即查询计 划通过Spark catalyst optimiser 进行优化。
比如下面一个例子:

为了说明查询优化,我们来看上图展示的人口数据分析的示例。图中构造了两个 DataFrame,将它们join 之后又做了一次filter操作。如果原封不动地执行这个执行计划,最 终的执行效率是不高的。因为join是一个代价较大的操作,也可能会产生一个较大的数据 集。如果我们能将filter下推到 join下方,先对DataFrame进行过滤,再join过滤后的较小 的结果集,便可以有效缩短执行时间。而Spark SQL的查询优化器正是这样做的。简而言之,逻辑查询计划优化就是一个利用基于关系代数的等价变换,将高成本的操作替换为低成本操作的过程。

6.3.2 DataSet
DataSet 是分布式数据集合。DataSet是Spark 1.6中添加的一个新抽象,是DataFrame 的一个扩展。它提供了RDD的优势(强类型,使用强大的lambda函数的能力)以及Spark SQL 优化执行引擎的优点。DataSet也可以使用功能性的转换(操作map,flatMap,filter 等等)。
- DataSet是DataFrame API 的一个扩展,是SparkSQL最新的数据抽象
- 用户友好的API风格,既具有类型安全检查也具有DataFrame的查询优化特性;
- 用样例类来对DataSet中定义数据的结构信息,样例类中每个属性的名称直接映射到 DataSet 中的字段名称;
- DataSet是强类型的。比如可以有DataSet[Car],DataSet[Person]。
- DataFrame是DataSet的特列,DataFrame=DataSet[Row] ,所以可以通过as方法将 DataFrame 转换为DataSet。Row 是一个类型,跟Car、Person这些的类型一样,所有的 表结构信息都用Row来表示。获取数据时需要指定顺序
6.4 DataFrame & DataSet的使用(spark-shell)
6.4.1 DataFrame
(1)创建DataFrame
a 从spark数据源中进行创建
- 查看Spark支持创建文件的数据源格式
scala> spark.read.
csv format jdbc json load option options orc parquet schema table text textFile
- 在spark的bin/data目录中创建user.json文件
|
{"username": "zhangsan", "age": "30"} {"username": "lisi", "age": "30"} {"username": "wangwu", "age": "40"} |
- 读取json文件创建DataFrame
scala> val df = spark.read.json("data/user.json")
df: org.apache.spark.sql.DataFrame = [age: bigint, username: string]
注意:如果从内存中获取数据,spark可以知道数据类型具体是什么。如果是数字,默认作 为Int 处理;但是从文件中读取的数字,不能确定是什么类型,所以用bigint接收,可以和 Long 类型转换,但是和Int不能进行转换
- 展示结果
scala> df.show
+---+--------+
|age|username|
+---+--------+
| 30|zhangsan|
| 30| lisi|
| 40| wangwu|
+---+--------+
b从RDD进行转换
在后续章节中讨论
c 从Hive Table 进行查询返回
在后续章节中讨论
(2)使用SQL语法操作DataFrame——通过spark.sql
SQL 语法风格是指我们查询数据的时候使用SQL语句来查询。由于DataFrame仅是一个二维数据集,并不是一个真正的表,因此使用这种风格的查询必须要有临时视图或者全局视图来辅助。
注意,由于我们是基于视图操作,所以只能进行查询,无法进行视图的修改。
a 临时视图的创建与查询
- 读取json文件创建DataFrame
scala> val df = spark.read.json("data/user.json")
df: org.apache.spark.sql.DataFrame = [age: bigint, username: string]
- 对DataFrame创建一个临时视图
scala> df.createOrReplaceTempView("userView")
- 使用SQL查询创建的临时视图(尽量使用createOrReplaceTempView,这样可以防止已经存在视图无法覆盖)
scala> val userViewStat = spark.sql("select * from userView")
userViewStat: org.apache.spark.sql.DataFrame = [age: string, username: string]
scala> userViewStat.show
+---+--------+
|age|username|
+---+--------+
| 30|zhangsan|
| 30| lisi|
| 40| wangwu|
+---+--------+
注意:普通临时表是Session范围内的,如果想应用范围内有效,可以使用全局临时表。
b 全局视图的创建与查询
- 我们使用createOrReplaceTempView临时视图时,在创建新的session后是无法查询到的
scala> spark.sql("select * from userView").show
+---+--------+
|age|username|
+---+--------+
| 30|zhangsan|
| 30| lisi|
| 40| wangwu|
+---+--------+
scala> spark.newSession.sql("select * from userView").show
25/09/25 17:18:01 WARN NativeIO: NativeIO.getStat error (3): 系统找不到指定的路径。
-- file path: tmp/hive
25/09/25 17:18:02 WARN HiveConf: HiveConf of name hive.stats.jdbc.timeout does not exist
25/09/25 17:18:02 WARN HiveConf: HiveConf of name hive.stats.retries.wait does not exist
- 若我们使用CreateOrReplaceGlobalTempView则可以在新session查询到旧session中的视图
scala> df.show
+---+--------+
|age|username|
+---+--------+
| 30|zhangsan|
| 30| lisi|
| 40| wangwu|
+---+--------+
scala> df.createOrReplaceTempView("emp")
scala> spark.sql("select * from global_temp.emp").show
+---+--------+
|age|username|
+---+--------+
| 30|zhangsan|
| 30| lisi|
| 40| wangwu|
+---+--------+
scala> spark.newSession.sql("select * from global_temp.emp").show
+---+--------+
|age|username|
+---+--------+
| 30|zhangsan|
| 30| lisi|
| 40| wangwu|
+---+--------+
注意:使用全局临时表时需要全路径访问,如:global_temp.emp,也就是在要访问的全局临时表前加上global_temp前缀。因为全局表会被注册到全局数据库global_temp中,不会被保存到默认数据库。
(3)DSL语法——通过DataFrame直接调用DSL(类似方法调用)
DataFrame提供一个特定领域语言(domain-specific language, DSL)去管理结构化的数据。 可以在 Scala, Java, Python 和 R 中使用 DSL,使用 DSL 语法风格不必去创建临时视图了。
- 创建一个DataFrame
scala> val df = spark.read.json("input/user.json")
df: org.apache.spark.sql.DataFrame = [age: string, username: string]
- DSL可以运行我们直接通过DataFrame去查询数据,不需要通过创建View & spark.sql去查询。比如我们可以调用df.select查询所有信息 或者 username。
#查询所有内容
scala> df.select("*").show
+---+--------+
|age|username|
+---+--------+
| 30|zhangsan|
| 30| lisi|
| 40| wangwu|
+---+--------+
#查询username列
scala> df.select("username").show
+--------+
|username|
+--------+
|zhangsan|
| lisi|
| wangwu|
+--------+
- 当我们要对某个查询结果进行计算时,需要查询列进行引用。DSL提供了两种引用方式,第一种是 $”字段” ,第二种是 ’字段 。
#查询age + 1的结果
# 使用 $”age” 进行引用
scala> df.select($"age" + 1).show
+---------+
|(age + 1)|
+---------+
| 31.0|
| 31.0|
| 41.0|
+---------+
# 使用 ’age 进行引用
scala> df.select('age + 1).show
+---------+
|(age + 1)|
+---------+
| 31.0|
| 31.0|
| 41.0|
+---------+
#查询username列 和 age + 1列的结果
scala> df.select('username, 'age + 1).show
+--------+---------+
|username|(age + 1)|
+--------+---------+
|zhangsan| 31.0|
| lisi| 31.0|
| wangwu| 41.0|
+--------+---------+
注意,若查询结果中包含引用字段,则所有字段都需要使用引用形式进行查询,否则无法查询到对应字段。
- 通过filter过滤DataFrame中的数据
#查询age > 30的数据
scala> df.filter('age > 30).show
+---+--------+
|age|username|
+---+--------+
| 40| wangwu|
+---+--------+
同理,这个也是一种对字段的计算,所以需要对所有查询字段进行引用标识。
- 通过groupBy进行分组
scala> df.groupBy("age").count.show
+---+-----+
|age|count|
+---+-----+
| 30| 2|
| 40| 1|
+---+-----+
(4)RDD 与 DataFrame的转换

RDD是无结构的数据,而DataFrame是有结构的数据,因此要将RDD转换成DataFrame就要赋予RDD对应的结构。
而DataFrame转换成RDD,由于DF本身就是结构化数据,所以直接转换即可。
a RDD To DataFrame——通过RDD.toDF(字段1, 字段2)
在IDEA中开发程序时,如果需要RDD与DF或者DS之间互相操作,那么需要引入 import spark.implicits._
这里的spark不是Scala中的包名,而是创建的sparkSession对象的变量名称,所以必 须先创建SparkSession对象再导入。这里的spark对象不能使用var声明,因为Scala只支持 val修饰的对象的引入。
spark-shell中无需导入,自动完成此操作。
- 比如我们将内存中创建的List[Int]数据集转换成字段为id的DataFrame,只需通过toDF(“id”) 指定生成的DF字段名即可:
scala> val idRDD = sc.makeRDD(List(1, 2, 3, 4))
idRDD: org.apache.spark.rdd.RDD[Int] = ParallelCollectionRDD[99] at makeRDD at <console>:24
//指定生成的DataFrame的字段名
scala> val df = idRDD.toDF("id")
df: org.apache.spark.sql.DataFrame = [id: int]
scala> df.show
+---+
| id|
+---+
| 1|
| 2|
| 3|
| 4|
+---+
- 实际生产中,更常见的是基于样例类将RDD转换成DF(这可以类比为设计一个数据库映射对象用于更方便将数据插入数据库),这个方法同时可以让RDD转换为DF & DS,是很方便的用法:
//首先定义一个样例类,用于映射RDD中的数据
scala> case class User(name: String, age: Int)
defined class User
//创建一个RDD,其中元素为一系列的(String, Int)元组
scala> val userRDD = sc.makeRDD(List(("zhangsan", 21), ("darren", 20)))
userRDD: org.apache.spark.rdd.RDD[(String, Int)] = ParallelCollectionRDD[103] at makeRDD at <console>:24
//将RDD中的元组元素依次遍历,转换成User对象,使数据结构化,然后转换为DF
scala> val userDF = userRDD.map({ case (name, age) => User(name, age) }).toDF
userDF: org.apache.spark.sql.DataFrame = [name: string, age: int]
scala> userDF.show
+--------+---+
| name|age|
+--------+---+
|zhangsan| 21|
| darren| 20|
+--------+---+
b DataFrame To RDD——通过DF.rdd
- DataFrame其实就是对RDD的封装,所以可以直接获取内部的RDD。
scala> val userRDD1 = userDF.rdd
userRDD1: org.apache.spark.rdd.RDD[org.apache.spark.sql.Row] = MapPartitionsRDD[112] at rdd at <console>:25
scala> userRDD1.foreach(println)
[zhangsan,21]
[darren,20]
注意:此时得到的RDD存储类型为Row
6.4.2 DataSet
DataSet是具有强类型的数据集合,需要提供对应的类型信息。
(1)创建DataSet(RDD转换成DataSet)
DataSet的创建有两种形式,RDD to DataSet 以及 序列化基本类型数据 to DataSet。由于序列化数据转换成DS基本不会使用,因此此处仅介绍RDD转换成DS的操作。
同样的,我们最好使用case class,将RDD数据提前结构化,这样就可以直接将结构化的RDD转换成DS:
scala> case class Person(name: String, age: Int)
defined class Person
scala> val personRDD = sc.makeRDD(List(Person("zhangsan", 30), Person("darren", 20)))
personRDD: org.apache.spark.rdd.RDD[Person] = ParallelCollectionRDD[113] at makeRDD at <console>:26
scala> val DS = personRDD.toD
toDF toDS toDebugString
scala> val DS = personRDD.toDS
DS: org.apache.spark.sql.Dataset[Person] = [name: string, age: int]
scala> DS.show
+--------+---+
| name|age|
+--------+---+
|zhangsan| 30|
| darren| 20|
+--------+---+
(2)DataSet 与 DataFrame的转换

首先我们需要理解DS和DF的具体区别,因为二者都是结构化数据集,看起来十分类似,容易混淆。
我们之前说,DataSet是具有强类型的数据集合。如何理解呢?
我们先看一下DataFrame。我们假设这个DataFrame是一个具有(name, age)两个字段的schema,则只需要满足(name, age)结构的数据都可以是DataFrame的成员,无论这样的数据是来自Person类还是Emp类。
而DataSet则要求(name, age)这样的结构化数据本身必须具有类型,比如其必须是Person类的(name, age)结构数据,才能够成为该DataSet的数据。其实,DataFrame是DataSet的特殊类,DataFrame每一个结构化数据都是Row类型,也就是说DataFrame = DataSet[Row]。
简单来说,就是DF的数据仅需要满足当前DF规定的字段即可;而DS的数据除了结构需要满足字段外,同时结构化数据类型也必须满足DS指定的类型。
a DataFrame To DataSet——通过DF.as[T]
我们提到,DS就是将DataFrame的结构化数据强制赋予类型,那么将DataFrame转换成DataSet也非常简单,直接通过DF.as[T]将DF中的结构化数据赋予样例类类型即可:
scala> val df = sc.makeRDD(List(("zhangsan", 11), ("darren", 20))).toDF("name", "age")
df: org.apache.spark.sql.DataFrame = [name: string, age: int]
scala> case class Person(name: String, age: Int)
defined class Person
scala> val ds = df.as[Person]
ds: org.apache.spark.sql.Dataset[Person] = [name: string, age: int]
scala> ds.show
+--------+---+
| name|age|
+--------+---+
|zhangsan| 11|
| darren| 20|
+--------+---+
b DataSet To DataFrame——通过DS.toDF
由于DS是DataFrame赋予类型,因此从DS中取出DF也十分简单,直接通过DS.toDF取出即可:
scala> val df = ds.toDF
df: org.apache.spark.sql.DataFrame = [name: string, age: int]
scala> df.show
+--------+---+
| name|age|
+--------+---+
|zhangsan| 11|
| darren| 20|
+--------+---+
(3)DataSet 与 RDD的直接转换

a RDD To DataSet——通过RDD.toDS
既然DataSet是有类型要求的结构化数据,那么我们直接在RDD中创建对应类型的结构化数据,然后将其构建成DataSet即可。注意,该方法需要有样例类进行DS的映射。
scala> val rdd = sc.makeRDD(List(Person("darren", 20), Person("lisi", 40)))
rdd: org.apache.spark.rdd.RDD[Person] = ParallelCollectionRDD[128] at makeRDD at <console>:26
scala> val ds = rdd.toDS
ds: org.apache.spark.sql.Dataset[Person] = [name: string, age: int]
scala> ds.show
+------+---+
| name|age|
+------+---+
|darren| 20|
| lisi| 40|
+------+---+
b DataSet To RDD——通过DS.rdd
从DS中取得RDD的方法与从DF取出RDD同理,由于二者都是对RDD的封装(DS是对DF的类型封装),因此直接通过DS.rdd取出即可。
scala> val rdd1 = ds.rdd
rdd1: org.apache.spark.rdd.RDD[Person] = MapPartitionsRDD[134] at rdd at <console>:25
scala> rdd1.collect().foreach(println)
Person(darren,20)
Person(lisi,40)
6.4.3 RDD & DataFrame & DataSet
在SparkSQL中Spark 为我们提供了两个新的抽象,分别是DataFrame和DataSet。他们和RDD有什么区别呢?
首先从版本的产生上来看:
- Spark1.0 => RDD
- Spark1.3 => DataFrame
- Spark1.6 => Dataset
如果同样的数据都给到这三个数据结构,他们分别计算之后,都会给出相同的结果。不同是的他们的执行效率和执行方式。在后期的Spark版本中,DataSet有可能会逐步取代RDD 和DataFrame 成为唯一的API接口。
(1)三者的共性
- RDD、DataFrame、DataSet全都是spark平台下的分布式弹性数据集,为处理超大型数 据提供便利;
- 三者都有惰性机制,在进行创建、转换,如map方法时,不会立即执行,只有在遇到 Action 如 foreach 时,三者才会开始遍历运算;
- 三者有许多共同的函数,如filter,排序等;
- 在对DataFrame和Dataset进行操作许多操作都需要这个包:import spark.implicits._(在 创建好SparkSession 对象后尽量直接导入,涉及隐式转换)
- 三者都会根据 Spark 的内存情况自动缓存运算,这样即使数据量很大,也不用担心会 内存溢出
- 三者都有partition的概念
- DataFrame和DataSet均可使用模式匹配获取各个字段的值和类型
(2)三者的区别
a RDD
- RDD一般和spark mllib同时使用
- RDD不支持sparksql操作
b DataFrame
- 与RDD和Dataset不同,DataFrame每一行的类型固定为Row,每一列的值没法直接访问,只有通过解析才能获取各个字段的值
- DataFrame与DataSet一般不与 spark mllib 同时使用
- DataFrame与DataSet均支持 SparkSQL 的操作,比如select,groupby之类,还能注册临时表/视窗,进行 sql 语句操作
- DataFrame与DataSet支持一些特别方便的保存方式,比如保存成csv,可以带上表 头,这样每一列的字段名一目了然(后面专门讲解)
c DataSet
- Dataset和DataFrame 拥有完全相同的成员函数,区别只是每一行的数据类型不同。 DataFrame 其实就是DataSet 的一个特例 type DataFrame = Dataset[Row]
- DataFrame也可以叫Dataset[Row],每一行的类型是Row,不解析,每一行究竟有哪 些字段,各个字段又是什么类型都无从得知,只能用上面提到的getAS方法或者共 性中的第七条提到的模式匹配拿出特定字段。而Dataset中,每一行是什么类型是 不一定的,在自定义了case class之后可以很自由的获得每一行的信息
(3)三者的转换

6.5 在IDEA操作SparkSQL
6.5.1 添加SparkSQL相关Maven依赖
<!--spark sql相关依赖-->
<dependency>
<groupId>org.apache.spark</groupId>
<artifactId>spark-sql_2.12</artifactId>
<version>3.0.0</version>
</dependency>
请与Spark、Scala版本保持一致。
6.5.2 创建SparkSQL上下文环境对象(类似Spark Core中的SparkContext)
//TODO 创建SparkSQL运行环境
val sparkConf = new SparkConf().setMaster("local[*]").setAppName("sparkSQL")
val spark = SparkSession.builder().config(sparkConf).getOrCreate()
//TODO 执行逻辑操作
//TODO 关闭环境
spark.close()
SparkSQL的上下文环境对象是SparkSession,这也是我们一直在spark-shell中使用的变量。我们与spark-shell中保持一致,在IDEA创建中也其为spark。
但是创建SparkSession的流程会与创建SparkContext有较大不同。这是SparkSession源码中官方推荐的创建SparkSession的方法,需要通过builder来进行创建:


原因为,SparkSession的构造器是私有的,我们无法直接通过new SparkSession的方式创建对象:

同时不难发现,SparkSession底层是已经封装了SparkContext的,因此在SparkSQL环境中同样可以对RDD等进行操作。
6.5.3 SparkSQL中操作DataFrame
(1)读取结构化文件生成DataFrame
//读取到DataFrame
val df: DataFrame = spark.read.json("datas/user.json")
df.show
/*
+---+--------+
|age|username|
+---+--------+
| 30|zhangsan|
| 20| lisi|
| 40| wangwu|
+---+--------+
*/
(2)使用SQL文操作DF
//使用SQL文操作DF
df.createOrReplaceTempView("user")
spark.sql("select * from user").show
/*
+---+--------+
|age|username|
+---+--------+
| 30|zhangsan|
| 20| lisi|
| 40| wangwu|
+---+--------+
*/
spark.sql("select avg(age) from user").show
/*
+--------+
|avg(age)|
+--------+
| 30.0|
+--------+
*/
(3)使用DSL操作DF
a 仅查询原始数据
在IDEA使用DSL操作DataFrame时,大体上与spark-shell命令差别不大:
//使用DSL操作DF
df.select("username").show
/*
+--------+
|username|
+--------+
|zhangsan|
| lisi|
| wangwu|
+--------+
*/
b 查询对原始数据进行转换的数据
但是,当涉及到我们要对DataFrame数据进行转换操作时(比如对age列进行 + 1的查询),IDEA与spark-shell就极为不同。
在spark-shell中,会自动导入SparkSQL上下文环境对象的隐式转换规则;而在IDEA中,这一步骤需要我们自己操作,也就是“import 当前SparkSQL上下文环境对象.implicits._“ ,在示例代码中,由于我们将当前SparkSQL上下文环境对象称作spark,因此引入语句是 “import spark.implicits._”。
//涉及到转换操作,需要引入转换规则(隐式参数)
import spark.implicits._
df.select($"age" + 1).show
/*
+---------+
|(age + 1)|
+---------+
| 31|
| 21|
| 41|
+---------+
*/
df.select('age + 1, 'username).show
/*
+---------+--------+
|(age + 1)|username|
+---------+--------+
| 31|zhangsan|
| 21| lisi|
| 41| wangwu|
+---------+--------+
*/
一般而言,建议将隐式转换的引入与SparkSession创建放到一起,避免遗忘。比如:
//创建SparkSQL环境
val sparkConf = new SparkConf().setMaster("local[*]").setAppName("sparkSQL")
val spark = SparkSession.builder().config(sparkConf).getOrCreate()
import spark.implicits._
6.5.4 SparkSQL中操作DataSet
(1)从Seq序列中创建DataSet
//从Seq序列中创建DS
val seq = Seq(1, 2, 3, 4)
val ds = seq.toDS()
ds.show()
/*
|value|
+-----+
| 1|
| 2|
| 3|
| 4|
+-----+
*/
(2)操作DataSet

我们查看DataFrame的源码则不难发现,DataFrame实际上是DataSet特定泛型的一个别名,因此DataFrame有的操作DataSet都有,直接参考6.5.3即可,此处不再过多赘述。
6.5.5 RDD – DataFrame – DataSet之间的相互转换
val sparkConf = new SparkConf().setMaster("local[*]").setAppName("SparkSQL")
val spark = SparkSession.builder.config(sparkConf).getOrCreate()
import spark.implicits._
//RDD <=> DataFrame
val rdd = spark.sparkContext.makeRDD(List(("darren", 20), ("lisi", 30), ("zhangsan", 40)))
//RDD => DataFrame
val df = rdd.toDF("name", "age")
df.show()
/*
+--------+---+
| name|age|
+--------+---+
| darren| 20|
| lisi| 30|
|zhangsan| 40|
+--------+---+
*/
//DataFrame => RDD
val rowRDD: RDD[Row] = df.rdd
rowRDD.foreach(println)
/*
[zhangsan,40]
[darren,20]
[lisi,30]
*/
//DataFrame <=> DataSet
//提供DataFrame
val df1 = spark
.sparkContext.makeRDD(List(("darren", 20), ("lisi", 30), ("zhangsan", 40)))
.toDF("name", "age")
//DataFrame => DataSet
val ds1 = df1.as[User]
ds1.show()
/*
+--------+---+
| name|age|
+--------+---+
| darren| 20|
| lisi| 30|
|zhangsan| 40|
+--------+---+
*/
//DataSet => DataFrame
val df2 = ds1.toDF
df2.show
//RDD <=> DataSet
//RDD => DataSet
//将rdd中的结构数据转换为业务类数据
val userRDD = spark
.sparkContext.makeRDD(List(("darren", 20), ("lisi", 30), ("zhangsan", 40)))
.map({ case (name, age) => User(name, age) })
//直接将具有特定类型的RDD转换为DS
val userDS = userRDD.toDS()
userDS.show()
/*
+--------+---+
| name|age|
+--------+---+
| darren| 20|
| lisi| 30|
|zhangsan| 40|
+--------+---+
*/
//DataSet => RDD
val userRDD1 = userDS.rdd
userRDD1.foreach(println)
/*
User(zhangsan,40)
User(darren,20)
User(lisi,30)
*/
spark.close()
(1)RDD ó DataFrame
由于SparkSession底层已经封装了SparkContext,所以我们只需要将该属性取出就可以使用对应的RDD操作方法:
//RDD <=> DataFrame
val rdd = spark.sparkContext.makeRDD(List(("darren", 20), ("lisi", 30), ("zhangsan", 40)))
a RDD => DataFrame
//RDD => DataFrame
val df = rdd.toDF("name", "age")
df.show()
/*
+--------+---+
| name|age|
+--------+---+
| darren| 20|
| lisi| 30|
|zhangsan| 40|
+--------+---+
*/
通过RDD.toDF方法,为转换后的DataFrame指定对应的列名即可。需要注意,RDD中的结构化元素需要能够与DataFrame结构一致。
b DataFrame => RDD
//DataFrame => RDD
val rowRDD: RDD[Row] = df.rdd
rowRDD.foreach(println)
/*
[zhangsan,40]
[darren,20]
[lisi,30]
*/
该操作实际上就是取出DF中的rdd属性。但是需要注意,我们原始RDD内的数据是tuple2类型的,从df中取出的RDD中的数据则是Row类型,二者有很大区别,需要注意区分!
(2)DataFrame ó DataSet
由于DataSet是具有特定类型的结构化数据(此处类型指的是业务类型,而非数据类型),因此我们提前提供样例类User:
//提供样例类
case class User(name: String, age: Int)
同时,我们提供一个与User具有相同结构的DataFrame:
val df1 = spark
.sparkContext.makeRDD(List(("darren", 20), ("lisi", 30), ("zhangsan", 40)))
.toDF("name", "age")
a DataFrame => DataSet
//DataFrame => DataSet
val ds1 = df1.as[User]
ds1.show()
/*
+--------+---+
| name|age|
+--------+---+
| darren| 20|
| lisi| 30|
|zhangsan| 40|
+--------+---+
*/
我们只需要为DataFrame中的结构化数据进行类型转换即可,即DF.as[T]。
b DataSet => DataFrame
//DataSet => DataFrame
val df2 = ds1.toDF
df2.show
/*
+--------+---+
| name|age|
+--------+---+
| darren| 20|
| lisi| 30|
|zhangsan| 40|
+--------+---+
*/
DS => DF更简单,去除DS的类型即可。
(3)RDD ó DataSet
我们同样使用case class User作为DataSet的业务类型:
//提供样例类
case class User(name: String, age: Int)
a RDD => DataSet
//RDD => DataSet
//将rdd中的结构数据转换为业务类数据
val userRDD = spark
.sparkContext.makeRDD(List(("darren", 20), ("lisi", 30), ("zhangsan", 40)))
.map({ case (name, age) => User(name, age) })
//直接将具有特定类型的RDD转换为DS
val userDS = userRDD.toDS()
userDS.show()
/*
+--------+---+
| name|age|
+--------+---+
| darren| 20|
| lisi| 30|
|zhangsan| 40|
+--------+---+
*/
RDD是可以直接转换为DataSet的,只要RDD中的数据都是特定类型的数据即可。在上述示例代码中,我们通过将RDD的结构化数据进行模式匹配,转换为User类,构成一个RDD[User],从而直接将RDD转换为DataSet。
b DataSet => RDD
//DataSet => RDD
val userRDD1 = userDS.rdd
userRDD1.foreach(println)
/*
User(zhangsan,40)
User(darren,20)
User(lisi,30)
*/
6.6 用户自定义函数
用户可以通过spark.udf功能添加自定义函数,实现自定义功能。
6.6.1 UDF函数
(1)UDF函数简介
UDF函数是较基础的用户自定义函数,用于对数据表中的每一条数据都进行转换,该操作不涉及分区shuffle,可以直接通过spark.udf.register注册相关函数名以及对应函数体。
(2)UDF函数使用——通过spark.udf.register直接注册
假设我们想要通过一个函数,为我们查询的数据加上我们想要的前缀,比如为查询到的name字段加上’NAME:’前缀,形成 NAME: name 。除了通过concat字符串拼接,我们也可以自己定义一个UDF函数,通过spark.udf实现此功能。具体实现如下:
val df = spark.read.json("datas/user.json")
df.createOrReplaceTempView("user")
//通过spark.udf注册一个自定义函数
spark.udf.register("prefixName", (name: String) => "Name: " + name)
spark.sql("select prefixName(username), age from user").show()
/*
+--------------------+---+
|prefixName(username)|age|
+--------------------+---+
| Name: zhangsan| 30|
| Name: lisi| 20|
| Name: wangwu| 40|
+--------------------+---+
*/
在使用udf函数时,步骤为:通过spark.udf.register为当前环境注册一个udf函数 -> 在sparkSQL文中使用这个udf函数。
spark.udf.register注册解析为:spark.udf.register(udf函数名, udf函数体)。
注册后的udf函数是通用的,只要sql中的形参满足函数体的形参要求即可。比如上述代码中,prefixName的函数体为 (name: String) => “Name: ” + name ,也就是为查询到的String类型字段加上Name: 前缀。
6.6.2 UDAF函数——用户定义聚合函数
(1)UDAF函数简介
UDAF函数即用户自定义聚合函数,与UDF的区别就在于,UDAF会对数据进行聚合,而非简单的逐条转换。因此,UDAF函数往往是涉及跨区域shuffle的。
使用UDAF可以允许我们在通过group by等sql语句时,在select语句中实现我们自定义的聚合操作。比如:“select udaf(age), name from t1 group by name“语句,正常的sql是不允许select出现不在group依据中的列名的,但是使用UDAF时则是允许的。
强类型的Dataset和弱类型的DataFrame都提供了相关的聚合函数, 如 count(), countDistinct(),avg(),max(),min()。
除此之外,用户可以设定自己的自定义聚合函数。通过继承UserDefinedAggregateFunction 来实现用户自定义弱类型聚合函数。从Spark3.0版本 后,UserDefinedAggregateFunction 已经不推荐使用了。可以统一采用强类型聚合函数 Aggregator。
需要注意的是,UDAF函数无法直接通过spark.udf.register注册,需要我们提前准备好函数逻辑,再通过spark.udf.register进行注册!
(2)实现UDAF函数的原理分析——以求平均值为例
现在我们想要通过一个UDF函数,实现对数据表Int类型字段的求平均值,比如对age字段求avg(age)。
其实现原理为:选出每一条数据中的age字段,进行age字段的累加以及计数,最后用age / cnt得出avg:

(3)通过弱类型实现UDAF求age字段均值——继承UserDefinedAggregateFunction类
package main.spark_sql
import org.apache.spark.SparkConf
import org.apache.spark.sql.{Row, SparkSession}
import org.apache.spark.sql.expressions.{MutableAggregationBuffer, UserDefinedAggregateFunction}
import org.apache.spark.sql.types.{DataType, LongType, StructField, StructType}
object SparkSQL_UDAF {
def main(args: Array[String]): Unit = {
val sparkConf = new SparkConf().setMaster("local[*]").setAppName("SparkSQL_UDF")
val spark = SparkSession.builder().config(sparkConf).getOrCreate()
import spark.implicits._
val df = spark.read.json("datas/user.json")
df.createOrReplaceTempView("user")
spark.udf.register("avgAge", new MyAvgUDAF)
spark.sql("select avgAge(age) from user").show
/*
+--------------+
|myavgudaf(age)|
+--------------+
| 30|
+--------------+
*/
spark.close()
}
/*
自定义聚合函数类
1、继承UserDefinedAggregateFunction并重写方法
2、重写其中的方法
*/
class MyAvgUDAF extends UserDefinedAggregateFunction{
//输入数据的结构
override def inputSchema: StructType = {
StructType(
Array(StructField("age", LongType))
)
}
//数据计算缓冲区的结构
override def bufferSchema: StructType = {
StructType(
Array(
StructField("total", LongType),
StructField("count", LongType)
)
)
}
//输出数据的数据类型
override def dataType: DataType = LongType
//函数的稳定性
override def deterministic: Boolean = true
//缓冲区初始化
override def initialize(buffer: MutableAggregationBuffer): Unit = {
//我们定义的缓冲区结构为:
//Array(
// StructField("total", LongType),
// StructField("count", LongType)
//)
buffer.update(0, 0L) //更新total处数据为0L
buffer.update(1, 0L) //更新count处数据为0L
}
//根据输入的值更新缓冲区数据
override def update(buffer: MutableAggregationBuffer, input: Row): Unit = {
buffer.update(0, buffer.getLong(0) + input.getLong(0))
buffer.update(1, buffer.getLong(1) + 1)
}
//合并不同缓冲区的数据结果
override def merge(buffer1: MutableAggregationBuffer, buffer2: Row): Unit = {
//更新到buffer1(scala默认逻辑,将第一个元素始终作为聚合初始值 & 迭代返回值)
buffer1.update(0, buffer1.getLong(0) + buffer2.getLong(0))
buffer1.update(1, buffer1.getLong(1) + buffer2.getLong(1))
}
//以最终的结果进行聚合计算,此处为平均值计算
override def evaluate(buffer: Row): Any = buffer.getLong(0) / buffer.getLong(1)
}
}
a 首先我们需要继承并重写UserDefinedAggregateFunctio类
/*
自定义聚合函数类
1、继承UserDefinedAggregateFunction并重写方法
2、重写其中的方法
*/
class MyAvgUDAF extends UserDefinedAggregateFunction{
//输入数据的结构
override def inputSchema: StructType = {
StructType(
Array(StructField("age", LongType))
)
}
//数据计算缓冲区的结构
override def bufferSchema: StructType = {
StructType(
Array(
StructField("total", LongType),
StructField("count", LongType)
)
)
}
//输出数据的数据类型
override def dataType: DataType = LongType
//函数的稳定性
override def deterministic: Boolean = true
//缓冲区初始化
override def initialize(buffer: MutableAggregationBuffer): Unit = {
//我们定义的缓冲区结构为:
//Array(
// StructField("total", LongType),
// StructField("count", LongType)
//)
buffer.update(0, 0L) //更新total处数据为0L
buffer.update(1, 0L) //更新count处数据为0L
}
//根据输入的值更新缓冲区数据
override def update(buffer: MutableAggregationBuffer, input: Row): Unit = {
buffer.update(0, buffer.getLong(0) + input.getLong(0))
buffer.update(1, buffer.getLong(1) + 1)
}
//合并不同缓冲区的数据结果
override def merge(buffer1: MutableAggregationBuffer, buffer2: Row): Unit = {
//更新到buffer1(scala默认逻辑,将第一个元素始终作为聚合初始值 & 迭代返回值)
buffer1.update(0, buffer1.getLong(0) + buffer2.getLong(0))
buffer1.update(1, buffer1.getLong(1) + buffer2.getLong(1))
}
//以最终的结果进行聚合计算,此处为平均值计算
override def evaluate(buffer: Row): Any = buffer.getLong(0) / buffer.getLong(1)
}
① inputSchema——指定输入数据的结构

这一步与②的bufferSchema一样,都是在指定数据的结构,以作为UDAF计算函数的实例结构参考。
我们深入讲解一下StructType和StructField的内容,以对后续内容有更详细的认知。
- 该方法的返回值是一个StructType类型的数据,因此我们先查看一下StructType的相关源码:

StructType是一个case class,因此允许我们可以不用通过new去创建对象。在创建StructType对象时,需要指定fields: Array[StructField],也就是说,StructType相当于是一系列StructField的封装。
简单来说,我们可以把StructType理解为一张表的数据结构定义,内部的每一个StructField是这张表中一个个字段的结构定义。
注意!这个StructType以及内部的StructField是元数据定义,仅定义了数据结构,但不包含数据具体的值。数据对应的值需要由实例去生成!
- 我们查看一下StructField源码,看一下我们具体要如何指定这个类型:

可以看到,在StructField的主构造器中,核心参数其实是name和dataType(另外两个参数都有默认值)。
参数具体解释为:
name:当前“字段”的名称
datatype:当前“字段“中数据的类型
- 对于StructType内的Array[StructType]的指定,spark官方给了我们这个示例供我们参考:

StructType(StructField1(name, dataType) :: StructField2 :: StructField3 :: Nil)
除了通过 :: 进行拼接构成Array,我们也可以通过我们最熟悉的Array工厂方法,直接去显示生成一个Array,即StructType(Array(StructType1, StructType2, …)),这也是我们在此处使用的实现手段。
- 有了上述的理解后,就可以解释我们实现的这个inputSchema方法重写了:

我们需要“一张表”作为输入数据的结构,即StructType数据类型。同时我们需要指定“输入数据表”具有的所有字段,也就是内部的Array[StructField]。
对于这张输入表,由于我们需要求平均值,因此只需要包含一个字段即可,因此StructType中的Array[StructField]仅包含一个元素。我们是对age字段求均值,此处我们将StructField.name设置为”age”;同时,age字段是num类型的,因此将StructField.dataType设置为LongType。
② bufferSchema——指定计算缓冲区的结构

我们在①中详细介绍了StructType以及其内部的StructField,因此此处我们仅介绍bufferSchema的设计思路。
bufferSchema是我们进行数据计算的缓冲区,而我们的数据计算是对输入数据求平均值。均值计算是 数据总和 / 数据个数,那么我们就需要两个参数,其一用于对输入数据进行求和,其二用于统计输入数据的个数。
因此,对于bufferSchema的StructType,需要封装两个StructField字段,其一用于保存总和,取名为”total”;第二个用于保存数据个数,起名为”count”。同时,由于两个参数均为num类型,因此dataType都指定为LongType类型。
③ datatype——指定输出数据(返回结果)的数据类型

输出数据是年龄均值,为num类型,为了简化结果,我们直接使用LongType整型作为输出数据类型。
④ deterministic——指定函数的稳定性

函数的稳定性简单理解就是:对于相同的输入,能否得出相同的结果。一般而言,对于不生成随机数的函数,相同输入总能够得到相同的输出。
对于求年龄均值,所有数据都是从TempView中取得,因此并不存在均值,所以函数稳定性为true,面对相同输入其输出均相同。
⑤ initialize——初始化数据计算缓冲区

initialize是对数据计算缓冲区实例MutableAggregationBuffer的更新,而MutableAggregationBuffer实例的生成与我们在bufferSchema中定义的数据计算缓冲区StructType结构是相对应的。
- 我们先查看一下MutableAggregationBuffer的相关源码:

可以看到,MutableAggregationBuffer是Row的子类,Row的部分源码如下:

不难发现,Row类中有一个关键的参数:StructType。在我们定义的UDAF中,该StructType实际上就是我们在bufferSchema中定义的缓冲区结构。
而MutableAggregationBuffe就会根据bufferSchema定义的缓冲区StructType生成包含各个字段的值的具体实例。
这些实例是与StructType中的StructField一一对应的,包括索引位置,因此我们通过索引就可以找到对应的StructField对应的值,这是我们进行update和获取字段值的基础。
- 接着我们就可以分析我们的initialize方法实现了:

基于我们定义的缓冲区数据StructType,由于其为Array(StructField-total, StructField-count),因此通过0索引获取的就是total字段对应的值,通过1索引获取的就是count字段对应的值。此时通过update第二个参数,指定更新后的值即可。
⑥ update——根据新的输入值更新数据计算缓冲区数据

我们已经在⑤中介绍了Row的结构,即Row同样是一个基于StructType生成的实例,因此该方法就不难理解了。
- 我们首先解释一下这个方法具有的两个参数:
- buffer:数据计算缓冲区的实例,是根据bufferSchema定义的StructType生成的。
- input:输入数据的实例,是根据inputSchema定义的StructType生成的。
- 于是这个方法的实现逻辑就很好理解了:
- input的StructType是一个Array(StructField-age),仅包含一个字段,因此通过0索引就可以取出该字段对应的实例值,即age。
- buffer的StructType是一个Array(StructField-total, StructField-count),所以0索引取出的是total,1索引取出的是count。
- 对于buffer的total更新,就是将buffer的0索引元素更新为原total 与新age之和,即:buffer.getLong(0)(buffer原有的total) + input.getLong(0)(input新输入的age)。
- 对于buffer的count更新,就是将buffer的1索引元素取出并 + 1,即:buffer.getLong(1) + 1。
⑦ merge——合并不同分区的缓冲区结果

对于两个缓冲区的合并逻辑实际上与update方法没什么区别,都是通过对应索引取出字段的实例值,然后进行求和。
此处需要注意的是,该方法返回值为Unit,我们更新方向是将另一个缓冲区buffer2的结果更新到当前缓冲区buffer1中。这是基于Scala集合底层的聚合思路实现的:永远将第一个集合作为初始值 & 迭代值,进行与其余集合的聚合运算。
⑧ evaluate——以最终的缓冲区结果进行聚合计算,得出最终结果

在最终缓冲区合并结束后,实例对应的值就是我们最后需要的值了。此时按照我们的业务逻辑进行计算即可。
在该UDAF中,我们是要求avg,因此计算逻辑为: total(buffer.getLong(0)) / count(buffer.getLong(1))。
b 在sparkSQL中通过spark.udf注册并使用UDAF
val df = spark.read.json("datas/user.json")
df.createOrReplaceTempView("user")
spark.udf.register("avgAge", new MyAvgUDAF)
spark.sql("select avgAge(age) from user").show
/*
+--------------+
|myavgudaf(age)|
+--------------+
| 30|
+--------------+
*/
具体的注册 & 使用流程与 6.6.1-(2)UDF函数使用 没有什么太大的差别,唯一注意的是我们此时要将我们实现的UDAF类作为函数体参数传给register进行注册。
通过spark.udf.register注册UDAF很好的体现了函数是一等公民的思想。
c 另一种StructType构造(通过StructField :: Nil)实现UDAF的参考

(4)通过强类型实现UDAF求均值——重点
在通过弱类型实现UDAF时,由于SQL传入的参数是一张表,因此没有类型的概念,我们只能通过将表映射成实体(StructType & StructField),然后通过列的序号对字段进行操作(比如通过buffer.update(0, 0L)来更新0索引的字段)。
这种操作无疑是十分麻烦的,同时在spark3.0中,已经不再推荐使用弱类型实现UDAF。我们可以发现我们通过弱类型实现的UDAF需要继承UserDefinedAggregateFunction,而该类在spark中已经被划去,不再推荐:


在spark源码中,有这样一段话:” Aggregator[IN, BUF, OUT] should now be registered as a UDF via the functions. udaf(agg) method.”,这意味着使用强类型实现UDAF是更加推荐的手段。
以下是通过强类型实现UDAF并使用的一个示例,我们会对其进行解析:
package main.spark_sql
import org.apache.spark.SparkConf
import org.apache.spark.sql.expressions.{Aggregator, MutableAggregationBuffer, UserDefinedAggregateFunction}
import org.apache.spark.sql.types.{DataType, LongType, StructField, StructType}
import org.apache.spark.sql.{Encoder, Encoders, Row, SparkSession, functions}
object SparkSQL_UDAF1 {
def main(args: Array[String]): Unit = {
val sparkConf = new SparkConf().setMaster("local[*]").setAppName("SparkSQL_UDF")
val spark = SparkSession.builder().config(sparkConf).getOrCreate()
val df = spark.read.json("datas/user.json")
df.createOrReplaceTempView("user")
//强类型实现的UDAF的特定注册方式
//functions.udaf用于将强类型的UDAF转换为弱类型的UDAF,以兼容register注册操作
spark.udf.register("avgAge", functions.udaf(new MyAvgUDAF))
spark.sql("select avgAge(age) from user").show
/*
+--------------+
|myavgudaf(age)|
+--------------+
| 30|
+--------------+
*/
spark.close()
}
//用于作为BUF参数的样例类
//计算均值的Buffer需要年龄总值 & 数据数量,因此样例类具有total & count两个参数分别对应
//同时,由于Buffer的参数需要动态计算,因此必须显示指定为var类型
case class Buf(var total: Long, var count: Long)
/*
自定义聚合函数类
1、继承Aggregator[IN, BUF, OUT]
泛型定义:
IN: 输入数据的类型。此处输入参数为age字段的值,因此为Long类型
BUF: 数据计算缓冲区的类型。我们通过自定义的样例类来实现
OUT: 输出数据的类型。此处输出参数为age字段的均值,同样为Long类型
2、重写方法
*/
class MyAvgUDAF extends Aggregator[Long, Buf, Long] {
//缓冲区初始化,即创建一个新的缓冲区,并且赋予初始值
override def zero: Buf = Buf(0, 0)
//根据输入数据更新缓冲区的数据
override def reduce(buf: Buf, in: Long): Buf = {
//更新缓冲区的属性值
buf.total = buf.total + in
buf.count = buf.count + 1
//将更新后的缓冲区对象返回
buf
}
//合并缓冲区,同样合并到第一个缓冲区
override def merge(buf1: Buf, buf2: Buf): Buf = {
buf1.total = buf1.total + buf2.total
buf1.count = buf1.count + buf2.count
buf1
}
//返回计算结果
override def finish(buf: Buf): Long = buf.total / buf.count
//指定编码,用于shuffle时的跨分区网络传输,写法较固定
//指定缓冲区的编码
override def bufferEncoder: Encoder[Buf] = Encoders.product
//指定输出的编码
override def outputEncoder: Encoder[Long] = Encoders.scalaLong
}
}
a 同样的,我们需要继承Aggregator[IN, BUF, OUT]来实现强类型的UDAF
/*
自定义聚合函数类
1、继承Aggregator[IN, BUF, OUT]
泛型定义:
IN: 输入数据的类型。此处输入参数为age字段的值,因此为Long类型
BUF: 数据计算缓冲区的类型。我们通过自定义的样例类来实现
OUT: 输出数据的类型。此处输出参数为age字段的均值,同样为Long类型
2、重写方法
*/
class MyAvgUDAF extends Aggregator[Long, Buf, Long] {
//缓冲区初始化,即创建一个新的缓冲区,并且赋予初始值
override def zero: Buf = Buf(0, 0)
//根据输入数据更新缓冲区的数据
override def reduce(buf: Buf, in: Long): Buf = {
//更新缓冲区的属性值
buf.total = buf.total + in
buf.count = buf.count + 1
//将更新后的缓冲区对象返回
buf
}
//合并缓冲区,同样合并到第一个缓冲区
override def merge(buf1: Buf, buf2: Buf): Buf = {
buf1.total = buf1.total + buf2.total
buf1.count = buf1.count + buf2.count
buf1
}
//返回计算结果
override def finish(buf: Buf): Long = buf.total / buf.count
//指定编码,用于shuffle时的跨分区网络传输,写法较固定
//指定缓冲区的编码
override def bufferEncoder: Encoder[Buf] = Encoders.product
//指定输出的编码
override def outputEncoder: Encoder[Long] = Encoders.scalaLong
}
① Aggregator类的形参泛型说明

该类具有三个参数,分别是IN、BUF、OUT。在通过弱类型实现UDAF后这三个参数其实不难理解,IN为输入参数的泛型,BUF为缓冲区参数的泛型,OUT为输出参数的泛型。
由于我们的UDAF是为实现自定义的对age求均值,因此IN为当前查询的age,也就是Long类型;OUT为age的均值,我们同样设定为Long类型。
我们需要着重讲解下BUF泛型。BUF是我们进行数据计算的缓冲区,对于当前的业务,我们需要在缓冲区中存储 年龄总和 & 数据条数,因此我们通过一个case class Buf来作为我们自定义的缓冲区泛型:

需要注意的是,在弱类型实现UDAF中,由于缓冲区同样是一张数据表,因此我们通过StructType指定字段;而在强类型实现UDAF中,缓冲区同样可以是一个类型,因此可以直接通过创建一个case class作为缓冲区的类型。
同时,由于case class的属性默认是val类型的,而缓冲区数据是随着我们数据输入而变化的,因此我们在主构造器中需要显示声明缓冲区的两个形参均为var类型,这样才能改变数据输入。
② zero——缓冲区初始化

缓冲区初始化顾名思义,就是创建一个新的缓冲区,并且赋予其初始值。对于我们求年龄均值的业务,就是创建缓冲区,并且对自定义缓冲区的total & count参数都指定初始值 0L 。
③ reduce——根据输入数据(查询到的新数据)更新缓冲区

该方法的返回值是一个新的buf,也就是更新后的新缓冲区。
对缓冲区的更新相比弱类型UDAF有了很大的简化,我们不再需要通过缓冲区的列序号 & 输入数据的列序号进行更新,只需要通过 buf.属性 & in 进行更新,最后返回新的缓冲区即可。
④ merge——跨分区shuffle,合并不同分区的缓冲区

同样的,直接通过 buf.属性 进行更新,在更新完成后返回更新后的缓冲区即可。
此处对于更新到哪个缓冲区并没有明确的要求,只需要保证数据更新到同一个缓冲区,并且返回的也是这个缓冲区即可。
⑤ finish——根据最终的缓冲区返回计算结果

缓冲区数据计算结束后,指定业务逻辑,返回对应的数据。
此处为求均值操作,因此直接通过 buf.total(年龄总和) / buf.count(数据条数) 作为返回值即可。
⑥ bufferEncoder & outputEncoder——指定缓冲区 & 输出数据编码格式

由于UDAF是需要进行跨分区shuffle的,也就是每个分区的 计算缓冲区& 输出数据 是需要通过网络进行传输的,因此我们需要指定 缓冲区& 输出数据 的编码,使得二者可以在网络传输时被序列化。
⑦ 注意事项
在reduce、merge中会对缓冲区进行更新,我们最好直接对传入缓冲区的属性进行更新,然后将更新后的缓冲区直接返回,避免创建新的Buf对象!

b 在sparkSQL中使用强类型实现的UDAF
//强类型实现的UDAF的特定注册方式
//functions.udaf用于将强类型的UDAF转换为弱类型的UDAF,以兼容register注册操作
spark.udf.register("avgAge", functions.udaf(new MyAvgUDAF))
spark.sql("select avgAge(age) from user").show
/*
+--------------+
|myavgudaf(age)|
+--------------+
| 30|
+--------------+
*/
由于spark.udf.register仅支持对弱类型实现的UDAF的注册,因此我们必须通过function.udaf将强类型实现的UDAF强转为弱类型的UDAF,才能够进行注册,即通过:spark.udf.register(UDAF名称, function.udaf(new 强类型实现的UDAF)) 进行使用。
c 强类型UDAF过程 与 自定义累加器 的区分
强类型的UDAF与自定义ACC的实现有很多相似之处:指定输入 / 输出类型,根据输入更新数据,合并两个计算结果等。
但不难发现,二者其实还是有非常大的区别的。
① 自定义累加器
在实现自定义累加器时,我们需要为累加器指定一个属性,用于存储每次累加后的结果。同时,每一次的add、merge,都是直接更新这个属性的状态,在返回时直接将属性返回(返回的属性一定是最新状态的属性)。
其特点可以总结为:
- ✅ 单实例状态:每个累加器只有一个 sum 属性实例
- ✅ 就地更新:add() 方法直接修改这个实例的状态
- ✅ 无返回值:操作不返回新对象,直接修改内部状态
② 强类型UDAF
与自定义累加器不同,我们并没有在强类型UDAF的实现类中指定属性作为缓冲区改变存储,而是在每一次aggregate、merge时,返回一个新的缓冲区实例。尽管这个实例通常是对输入形参的修改,但是我们并不通过修改一个全局的属性来实现分区的aggregate,也就是说旧的缓冲区在旧的状态中是并没有被更改的。
其特点可以总结为:
- ✅ 无可变状态:AverageBuffer 是不可变的case class
- ✅ 返回新实例:每次操作都创建并返回新的缓冲区实例
- ✅ 函数式风格:不修改输入参数,纯函数式操作
③ 二者的实际影响对比
内存使用:
- 累加器:内存使用固定(单个实例)
- UDAF:可能创建多个中间实例,但Spark会优化重用
线程安全:
- 累加器:内部处理并发访问
- UDAF:天然线程安全(不可变对象)
调试难度:
- 累加器:状态变化轨迹难以追踪
- UDAF:每次转换都产生新状态,更容易推理
6.7 数据读取与保存
6.7.1 通用的读取 & 保存方法
(1)spark.read.load(“文件”)——通用读取
a 基础读取——读取parquet文件
在spark中,spark.read.load提供了通用的读取文件的方法,默认读取文件必须是parquet格式:
scala> val df = spark.read.load("/opt/module/spark-local/examples/src/main/resources/users.parquet")
df: org.apache.spark.sql.DataFrame = [name: string, favorite_color: string ... 1 more field]
scala> df.show
+------+--------------+----------------+
| name|favorite_color|favorite_numbers|
+------+--------------+----------------+
|Alyssa| null| [3, 9, 15, 20]|
| Ben| red| []|
+------+--------------+----------------+
b 读取其余格式的文件
若我们想要读取特定格式的文件,需要通过format指定转换的文件格式。比如我们想要通过spark.read.load读取json文件,就要以这种格式读取:spark.read.format(“json”).load(“你的文件路径”):
scala> val df = spark.read.format("json").load("data/user.json")
df: org.apache.spark.sql.DataFrame = [age: bigint, username: string]
scala> df.show
+---+--------+
|age|username|
+---+--------+
| 30|zhangsan|
| 20| lisi|
| 40| wangwu|
+---+--------+
(2)DF.write.save(“路径”)——通用保存
a 基础保存
该方法默认会将df中的数据保存为parquet文件,若使用相对路径则默认在spark安装目录下新建一个文件夹保存:
scala> df.write.save("output1")

同样的,若要保存为指定格式的文件,也需要通过format进行格式转换,比如df.write.format(“path”).save就是将df中的数据保存为json格式文件。
b 通过mode指定数据的保存方式
同时,保存操作可以使用 SaveMode, 用来指明如何处理数据,使用mode()方法来设置。有一点很重要: 这些 SaveMode 都是没有加锁的, 也不是原子操作。SaveMode具体如下:

比如我们想要在已有文件夹output中将新的df文件追加进去,就可以通过df.write.save + SaveMode实现:
scala> df.write.mode("append").save("output")
6.7.2 操作JSON / CSV文件
(1)操作JSON
Spark SQL 能够自动推测JSON数据集的结构,并将它加载为一个Dataset[Row]. 可以 通过SparkSession.read.json()去加载 JSON 文件。
注意:Spark读取的JSON文件不是传统的JSON文件,每一行都应该是一个JSON串。格式如下:
{"name":"Michael"}
{"name":"Andy", "age":30}
[{"name":"Justin", "age":19},{"name":"Justin", "age":19}]
除了使用spark.read.format(“_”).load来读取其余格式(除parquet)文件外,sparkSQL提供了更简洁的读取,可以通过spark.read.去查看:

我们可以通过spark.read.json直接读取json文件:
scala> spark.read.
csv jdbc load options parquet table textFile
format json option orc schema text
scala> val df = spark.read.json("data/user.json")
df: org.apache.spark.sql.DataFrame = [age: bigint, username: string]
但是若我们想要保存成json文件,则只能通过df.write.formate(“json”).save来进行保存,没有更加快捷的方式。
(2)操作CSV文件
我们读取csv格式文件时,需要通过设置一些配置项来配置读取的csv文件,比如:
scala> val df = spark.read.format("csv").option("seq", "-").option("inferSchema", "true").option("header", "true").load("/opt/module/spark-local/examples/src/main/resources/people.csv")
df: org.apache.spark.sql.DataFrame = [name;age;job: string]
scala> df.show
+------------------+
| name;age;job|
+------------------+
|Jorge;30;Developer|
| Bob;32;Developer|
+------------------+
我们在其中进行的配置说明如下:
val df = spark.read
.format("csv") // 指定数据源格式为CSV
.option("seq", "-") // 设置字段分隔符为短横线(-)
.option("inferSchema", "true") // 自动推断列的数据类型
.option("header", "true") // 将第一行作为列名
.load("/opt/module/spark-local/examples/src/main/resources/people.csv")
6.7.3 操作MySQL
Spark SQL可以通过JDBC从关系型数据库中读取数据的方式创建DataFrame,通过对 DataFrame一系列的计算后,还可以将数据再写回关系型数据库中。
如果使用spark-shell操作,可在启动shell时指定相关的数据库驱动路径或者将相关的数据库驱动放到spark的类 路径下:
bin/spark-shell
--jars mysql-connector-java-5.1.27-bin.jar
我们这里只演示在Idea中通过JDBC对Mysql进行操作。
(1)引入依赖
<!--mysql jdbc驱动相关依赖-->
<dependency>
<groupId>mysql</groupId>
<artifactId>mysql-connector-java</artifactId>
<version>8.0.33</version>
</dependency>
注意,该依赖版本需要与当前MySQL的版本保持一致。
(2)读取数据
我们在读取数据时,依旧采取通用的spark.read.load方式。由于读取的是jdbc文件而不是parquet,所以需要通过format指定数据源格式:
//TODO 通过JDBC读取MySQL数据
//将MySQL数据表数据读取到DataFrame中
val df = spark.read
.format("jdbc") //指定读取数据格式为jdbc
.option("url", "jdbc:mysql://localhost:3306/spark-sql") //要读取的MySQL数据库
.option("drive", "com.mysql.jdbc.driver") //指定jdbc驱动
.option("user", "root") //指定用户名
.option("password", "D200504193010") //指定密码
.option("dbtable", "user") //指定要读取的数据表
.load
df.show
//+---+--------+---+
//| id| name|age|
//+---+--------+---+
//| 1|zhangsan| 20|
//| 2| darren| 20|
//| 3| lisi| 40|
//| 4| wangwu| 30|
//+---+--------+---+
整体上与通过SparkCore+JDBC读取数据十分类似,但是更为优雅。
在SparkCore + JDBC的读取中,我们需要通过Connection、Statement、ResultSet三个JDBC变量进行操作,比如如下代码就是一个通过原始SC + JDBC的读取示例:
// 传统 JDBC 方式
import java.sql.{Connection, DriverManager, ResultSet, Statement}
Class.forName("com.mysql.cj.jdbc.Driver")
val connection: Connection = DriverManager.getConnection(
"jdbc:mysql://localhost:3306/test", "username", "password"
)
val statement: Statement = connection.createStatement()
val resultSet: ResultSet = statement.executeQuery("SELECT * FROM users")
while (resultSet.next()) {
val id = resultSet.getInt("id")
val name = resultSet.getString("name")
// 手动处理每一行数据
}
resultSet.close()
statement.close()
connection.close()
但是SparkSQL中,我们只需要通过options指定配置即可,因为SparkSQL底层对sparkCore + JDBC的原始读取进行了封装。
(3)写入数据
同样的,我们使用spark.write.save保存数据。我们想要将数据保存为一个新的MySQL数据表,因此也需要通过format指定保存数据为jdbc格式。
下面的示例是在6.7.3-(2)的读取基础上进行保存的:
//TODO 将读取到的数据保存到MySQL数据库
df.write
.format("jdbc") //指定保存数据格式为jdbc
.option("url", "jdbc:mysql://localhost:3306/spark-sql") //MySQL数据库
.option("drive", "com.mysql.jdbc.driver") //指定jdbc驱动
.option("user", "root") //指定用户名
.option("password", "D200504193010") //指定密码
.option("dbtable", "user1") //指定要保存的数据表
.mode(SaveMode.Append) //指定保存方式为追加
.save()
保存过程可能会有延迟。这是保存后的结果:

6.7.4 操作Hive
Apache Hive 是 Hadoop 上的 SQL 引擎,Spark SQL编译时可以包含 Hive 支持,也 可以不包含。包含 Hive 支持的 Spark SQL 可以支持 Hive 表访问、UDF (用户自定义函数) 以及 Hive 查询语言(HiveQL/HQL)等。需要强调的一点是,如果要在 Spark SQL 中包含 Hive 的库,并不需要事先安装 Hive。一般来说,最好还是在编译Spark SQL时引入Hive 支持,这样就可以使用这些特性了。如果你下载的是二进制版本的 Spark,它应该已经在编 译时添加了 Hive 支持。
若要把 Spark SQL 连接到一个部署好的 Hive 上,你必须把 hive-site.xml 复制到 Spark 的配置文件目录中($SPARK_HOME/conf)。即使没有部署好 Hive,Spark SQL 也可以 运行。 需要注意的是,如果你没有部署好Hive,Spark SQL 会在当前的工作目录中创建出 自己的 Hive 元数据仓库,叫作 metastore_db。此外,如果你尝试使用 HiveQL 中的 CREATE TABLE (并非 CREATE EXTERNAL TABLE)语句来创建表,这些表会被放在你默 认的文件系统中的 /user/hive/warehouse 目录中(如果你的 classpath 中有配好的 hdfs-site.xml,默认的文件系统就是 HDFS,否则就是本地文件系统)。 spark-shell 默认是 Hive 支持的;代码中是默认不支持的,需要手动指定(加一个参数即可)。
(1)操作SparkSQL内置Hive
接下来我们演示操作SparkSQL的内置Hive。在没有提前部署hive的情况下,流程为:创建Hive的元数据仓库metastore_db ——> 创建hive数据表存放位置spark-warehouse。
这是spark-local模式下的spark文件夹:

我们尝试通过spark.sql去查询hive数据表:
scala> spark.sql("show tables").show
#未部署hive时内置hive的元数据仓库创建…

+--------+---------+-----------+
|database|tableName|isTemporary|
+--------+---------+-----------+
+--------+---------+-----------+
我们会发现,由于hive没有提前部署,所以spark在运行这个命令时创建了一个hive文件夹,用于存储hive的元数据。我们刷新spark-local下的文件夹会发现多了一个文件夹metastore_db,这就是spark内置的hive元数据仓库:

我们通过spark.sql读取data/user.json数据,并且创建为临时表,然后再次查看数据表情况:
scala> val df = spark.read.json("data/user.json")
df: org.apache.spark.sql.DataFrame = [age: bigint, username: string]
scala> df.createOrReplace
createOrReplaceGlobalTempView createOrReplaceTempView
scala> df.createOrReplaceTempView("user")
scala> spark.sql("show tables").show
+--------+---------+-----------+
|database|tableName|isTemporary|
+--------+---------+-----------+
| | user| true|
+--------+---------+-----------+
可以发现,内部多了一张临时表。
接着,我们尝试通过sparkSQL执行hiveQL,将linux本地文件上传到我们创建的hive数据表文件夹中:
scala> spark.sql("create table test(id Int)")
#创建hive数据表…

res3: org.apache.spark.sql.DataFrame = []
#通过sparkSQL执行HiveQL上传本地文件到新建的表中
scala> spark.sql("load data local inpath 'data/id.txt' into table test")
res4: org.apache.spark.sql.DataFrame = []
scala> spark.sql("show tables").show
+--------+---------+-----------+
|database|tableName|isTemporary|
+--------+---------+-----------+
| default| test| false|
| | user| true|
+--------+---------+-----------+
scala> spark.sql("select * from test").show
+---+
| id|
+---+
| 1|
| 2|
| 3|
| 4|
| 5|
| 6|
+---+
不难发现,此时多了一张非临时表test,我们可以直接通过sparkSQL读取hive中的test表数据。同时我们可以查看spark-local目录,新创建了spark-warehouse文件夹,这就是hive数据表存储的位置:

同时,该文件夹中也有我们刚刚创建的数据表test:

文件夹内部有相关数据(我们创建的是内部表,数据存放在数据表文件夹下):

(2)操作外部已部署好的Hive——通过spark-shell
如果想连接外部已经部署好的Hive,需要通过以下几个步骤:
- Spark要接管Hive需要把hive-site.xml拷贝到conf/目录下
- 把Mysql的驱动copy到jars/目录下
- 如果访问不到hdfs,则需要把core-site.xml和hdfs-site.xml拷贝到conf/目录下
- 重启spark-shell
a 将部署好的hive与spark建立连接
① 将hive-site.xml拷贝到当前spark的conf/目录下
[root@node1 spark-local-outerHiveConnect]# su - hadoop
Last login: Wed Aug 27 18:34:01 CST 2025 on pts/1
[hadoop@node1 ~]$ cd /export/server/
[hadoop@node1 server]$ cd apache-hive-3.1.3-bin/
#拷贝文件
[hadoop@node1 apache-hive-3.1.3-bin]$ sudo cp conf/hive-site.xml /opt/module/spark-local-outerHiveConnect/conf/
[sudo] password for hadoop:
[hadoop@node1 apache-hive-3.1.3-bin]$
在我的本地中,hadoop & hive相关配置被授权给了hadoop用户,因此需要通过hadoop用户 & sudo命令拷贝到root用户的spark中。
② 把MySQL驱动拷贝到当前spark的jars/目录下
将资料中的驱动jar包上传即可:

b 重启spark-shell,连接hive
在重启spark-shell连接hive前,需要确保hadoop & hive都已经启动!
scala> spark.sql("show tables").show
#连接外部Hive…
25/10/02 16:36:04 WARN HiveConf: HiveConf of name hive.metastore.event.db.notification.api.auth does not exist
+--------+---------+-----------+
|database|tableName|isTemporary|
+--------+---------+-----------+
| default| test| false|
+--------+---------+-----------+
可以发现,并没有提示找不到Hive的元数据仓库,可以直接查询到hive的数据库(因为我的配置hive只有一个数据库表,因此仅显示了默认数据库的test表)。
同时,spark文件夹下也并未出现metastore_db(内置hive源数据仓库) & spark-warehouse(内置hive数据表存放文件夹)两个文件夹:

这说明,连接外置部署好的hive十分成功!
(3)操作外部已部署好的Hive——通过IDEA在代码中操作
a 导入依赖
注意,必须包含MySQL驱动!
<!--操作hive相关依赖-->
<dependency>
<groupId>org.apache.spark</groupId>
<artifactId>spark-hive_2.12</artifactId>
<version>3.0.0</version>
</dependency>
<!--mysql jdbc驱动相关依赖-->
<dependency>
<groupId>mysql</groupId>
<artifactId>mysql-connector-java</artifactId>
<version>8.0.33</version>
</dependency>
b 将hive-site.xml拷贝到当前项目的resources文件夹下,作为连接hive的配置文件

c 在sparkSQL中用代码访问并且操作hive
//TODO 通过SparkSQL连接外置Hive
val spark = SparkSession.builder()
.enableHiveSupport() //hive支持
.config(sparkConf).getOrCreate()
import spark.implicits._
spark.sql("show tables").show
(4)通过spark sql cli操作hive
Spark SQL CLI 可以很方便的在本地运行Hive元数据服务以及从命令行执行查询任务。在 Spark 目录下执行如下命令启动Spark SQL CLI,直接执行SQL语句,类似一Hive窗口:
bin/spark-sql

(5)通过beeline操作hive
Spark Thrift Server 是 Spark 社区基于 HiveServer2 实现的一个Thrift 服务。旨在无缝兼容 HiveServer2。因为 Spark Thrift Server 的接口和协议都和 HiveServer2 完全一致,因此我们部 署好Spark Thrift Server 后,可以直接使用hive的 beeline 访问Spark Thrift Server 执行相关 语句。Spark Thrift Server 的目的也只是取代HiveServer2,因此它依旧可以和Hive Metastore 进行交互,获取到hive的元数据。
如果想连接Thrift Server,需要通过以下几个步骤:
- Spark要接管Hive需要把hive-site.xml拷贝到conf/目录下
- 把Mysql的驱动copy到jars/目录下
- 如果访问不到hdfs,则需要把core-site.xml和hdfs-site.xml拷贝到conf/目录下
- 启动Thrift Server:
我们通过以下命令启动Thrift Server:
sbin/start-thriftserver.sh
- 然后就可以使用beeline连接Thrift Server了:
bin/beeline -u jdbc:hive2://node1:10000 -n root
注意,beeline的hive连接地址需要改成hive-site.xml中配置的地址!

6.8 SparkSQL案例实操
注:由于Hive无法使用,此处使用MySQL代替Hive来进行sparkSQL的读取。
6.8.1 数据准备
(1)建立数据库 & 数据表
# 创建数据库
create database if not exists spark_sql_practice;
#创建数据表
USE spark_sql_practice;
-- 用户行为表
CREATE TABLE user_visit_action (
date VARCHAR(20),
user_id BIGINT,
session_id VARCHAR(50),
page_id BIGINT,
action_time VARCHAR(50),
search_keyword VARCHAR(100),
click_category_id BIGINT,
click_product_id BIGINT,
order_category_ids VARCHAR(200),
order_product_ids VARCHAR(200),
pay_category_ids VARCHAR(200),
pay_product_ids VARCHAR(200),
city_id BIGINT
);
-- 产品信息表
CREATE TABLE product_info (
product_id BIGINT,
product_name VARCHAR(200),
extend_info VARCHAR(500)
);
-- 城市信息表
CREATE TABLE city_info (
city_id BIGINT,
city_name VARCHAR(50),
area VARCHAR(50)
);
(2)将数据读取到MySQL中
a 数据文件位置查看
这是我的数据文件所在的目录:

b 开启MySQL的local-infile服务
由于我们需要从本地读取文件,而MySQL对此项默认是关闭的,所以我们需要显示将其打开。注意,该选项需要同时开启server & client两端,也就是在MySQL配置文件中 & IDEA的MySQL连接中 都需要配置!
- 首先打开MySQL的配置文件(通常位于安装目录MySQL Server XX)的my.ini文件,在[mysqld]模块下,添加”local-infile = 1”:

- 然后我们可以通过MySQL CLI ,使用”show global variables like ‘local_infile’”语句查看是否成功添加:

MySQL CLI可以直接在MySQL的安装目录(MySQL Server XX)的bin文件夹下,通过cmd进入命令控制行,然后输入 “mysql -u 你的用户名 -p” 进入。
- 最后在IDEA的数据库连接池的properties中的Advanced处,指定allowLoadLocalInfile项为true即可:

c 读取文件到MySQL中
通过以下语句将数据载入MySQL:
-- 在 MySQL 中执行数据导入
# 载入user_visit_action数据
LOAD DATA LOCAL INFILE 'D:/develop/dataScience_learn/spark_learn/IDEA_PROJECT/Spark_IDEA/spark_core/src/main/spark_sql/practice/data/user_visit_action.txt'
INTO TABLE user_visit_action
FIELDS TERMINATED BY '\t'
LINES TERMINATED BY '\n';
# 载入product_info数据
LOAD DATA LOCAL INFILE 'D:/develop/dataScience_learn/spark_learn/IDEA_PROJECT/Spark_IDEA/spark_core/src/main/spark_sql/practice/data/product_info.txt'
INTO TABLE product_info
FIELDS TERMINATED BY '\t'
LINES TERMINATED BY '\n';
# 载入city_info数据
LOAD DATA LOCAL INFILE 'D:/develop/dataScience_learn/spark_learn/IDEA_PROJECT/Spark_IDEA/spark_core/src/main/spark_sql/practice/data/city_info.txt'
INTO TABLE city_info
FIELDS TERMINATED BY '\t'
LINES TERMINATED BY '\n';
其实我们会发现,整体来说与Hive的加载文件差异并不大。但是我们会发现,hive是在建表时就指定表文件的fields & lines的分隔符,而MySQL则是在载入数据时指定。
这是因为,在Hive中,表只是对实际文件的映射,因此表必须与文件格式保持一致。我们可以这样理解:Hive中的表格式是为了适配要存储在其中的文件的,我们其实是基于数据文件的格式指定表的格式。
而在MySQL中,表是一种确定的结构,数据是实际存储在表中的,而非作为映射,所以要求数据文件必须遵循一定的格式(比如列按制表符’\t’,行按换行’\n’分隔),才能对文件进行分隔,依次放到表中对应字段和数据行中。我们可以这样理解:MySQL中的表属于一种硬性要求,数据文件如果想要读取到表中,必须对数据文件的字段分隔、数据行分隔进行明确的指定。
d MySQL数据查看
我们通过SparkSQL测试一下数据是否成功加载到MySQL中:
package main.spark_sql.practice
import org.apache.spark.SparkConf
import org.apache.spark.sql.SparkSession
object code {
def main(args: Array[String]): Unit = {
val sparkConf = new SparkConf().setMaster("local[*]").setAppName("spark-sql-practice")
val spark = SparkSession.builder().config(sparkConf).getOrCreate()
// TODO 读取MySQL中的数据文件
val df = spark.read
.format("jdbc") //指定数据源为jdbc文件
.option("url", "jdbc:mysql://localhost:3306/spark_sql_practice") //要读取的MySQL数据库
.option("drive", "com.mysql.jdbc.driver") //指定jdbc驱动
.option("user", "root") //指定用户名
.option("password", "D200504193010") //指定密码
.option("dbtable", "city_info") //指定要读取的数据表
.load
df.show()
spark.close()
}
}
结果输出为:

显然,成功了~
6.8.2 需求介绍

6.8.3 需求分析

6.8.4 代码实现
(1)UDAF之前的代码实现

这部分需求的主要难点在于如何对每个地区的商品按照点击次数降序并且取出每个地区的前三,需要使用到窗口函数。
对于这类对数据的查询 & 筛选,此处提供两种方案:通过sql文实现 & 纯DataFrameAPI。
对于sql文的实现,由于连接、分组等操作都是基础操作,此处略过不提,重点介绍窗口函数的使用;对于DataFrameAPI,由于操作较为不同,我们会逐步介绍。
a 通过sql文实现
① 查询sql文展示
# 基础的数据表(点击数据与城市、产品的连接查询结果)
with basicTable as (
select
u.*, p.product_name, c.city_name, c.area
from
user_visit_action u, product_info p, city_info c
where
u.city_id = c.city_id
and u.click_product_id = p.product_id
and u.click_product_id != -1
),
# 对基础的数据表按照地区 & 商品id分组,统计每个商品在每个地区的总点击次数
groupTable as (
select
area, product_name, count(*) cnt
from
basicTable
group by area, product_name
),
# 按照地区分组,对每个地区的count降序排列,对每个分区数据分别进行排名
# 使用窗口函数实现该需求
orderTable as (
select
*,
rank() over ( partition by area order by cnt desc ) rank_num
from
groupTable
)
# 最终结果的查询(由于rank_num在where执行时还未创建,所以无法直接在上面的表中进行rank_num <= 3的筛选
select
*
from
orderTable
where
rank_num <= 3;
我们通过with语句建立临时表代替直接使用嵌套子查询,这样更容易理解。
② 窗口函数使用解析

在MySQL的sql文中,窗口函数为rank() over(),我们可以在over中指定分区标准 以及 分区数据排序依据。我们使用的语句为:
rank() over ( partition by area order by cnt desc ) rank_num
即:为groupTable表新增组内排名列,通过groupTable的area字段进行分组,对每一组都按照cnt字段值进行降序排列,并最终将该列命名为rank_num。
需要注意的是,对于rank_num的筛选必须新建一个查询。因为MySQL的实际执行中,select 语句会在 where之后执行,也就是说where只能对groupTable原有的列进行筛选,新增的rank列在where执行时还未存在。
③ 完整的sparkSQL代码
package main.spark_sql.practice
import org.apache.spark.SparkConf
import org.apache.spark.sql.SparkSession
//基于sql文实现需求
object CodeWithSQLText {
def main(args: Array[String]): Unit = {
val sparkConf = new SparkConf().setMaster("local[*]").setAppName("spark-sql-practice")
val spark = SparkSession.builder().config(sparkConf).getOrCreate()
//点击表
val userDF = spark.read.format("jdbc")
.option("url", "jdbc:mysql://localhost:3306/spark_sql_practice") //要读取的MySQL数据库
.option("drive", "com.mysql.jdbc.driver") //指定jdbc驱动
.option("user", "root") //指定用户名
.option("password", "D200504193010") //指定密码
.option("dbtable", "user_visit_action") //指定要读取的数据表
.load
//产品表
val productDF = spark.read.format("jdbc")
.option("url", "jdbc:mysql://localhost:3306/spark_sql_practice") //要读取的MySQL数据库
.option("drive", "com.mysql.jdbc.driver") //指定jdbc驱动
.option("user", "root") //指定用户名
.option("password", "D200504193010") //指定密码
.option("dbtable", "product_info") //指定要读取的数据表
.load
//城市表
val cityDF = spark.read.format("jdbc")
.option("url", "jdbc:mysql://localhost:3306/spark_sql_practice") //要读取的MySQL数据库
.option("drive", "com.mysql.jdbc.driver") //指定jdbc驱动
.option("user", "root") //指定用户名
.option("password", "D200504193010") //指定密码
.option("dbtable", "city_info") //指定要读取的数据表
.load
//直接通过sql文查询
//首先创建三个临时表,用于sql文使用
userDF.createOrReplaceTempView("user_visit_action")
productDF.createOrReplaceTempView("product_info")
cityDF.createOrReplaceTempView("city_info")
spark.sql(
"""
|with basicTable as (
| select
| u.*, p.product_name, c.city_name, c.area
| from
| user_visit_action u, product_info p, city_info c
| where
| u.city_id = c.city_id
| and u.click_product_id = p.product_id
| and u.click_product_id != -1
|),
|groupTable as (
| select
| area, product_name, count(*) cnt
| from
| basicTable
| group by area, product_name
|),
|orderTable as (
| select
| *,
| rank() over ( partition by area order by cnt desc ) rank_num
| from
| groupTable
|)
|select
| *
|from
| orderTable
|where
| rank_num <= 3;
|""".stripMargin).show()
//TODO 城市备注需要使用自定义UDAF函数
spark.close()
}
}
需要注意,由于MySQL是spark外部的数据库,因此无法通过”use database”等语句操作(这些语句是操作spark内置数据库的)。所以我们必须将数据先通过读取jdbc载入到sparkDF中,然后创建临时表,基于临时表查询。
同时,在sparkSQL的sql文中不能有MySQL形式的注释,若出现则会导致sparkSQL无法解析。
b 通过纯DataFrameAPI实现
同样的,我们必须先将MySQL数据通过read.format(“jdbc”).load导入到spark中,才能进行后续操作:

① 连接三张表的数据,获取完整的数据(只有点击)

对于点击数据,要求为userDF的click_product_id > -1,所以通过userDF作为过滤 & 连接操作的对象。
在DataFrameAPI中,可以直接通过join进行两个DF的连接。
join操作的参数为:
-
- 要与调用者进行连接的DataFrame
- 两个DataFrame的连接依据(基于哪个列连接)
- 连接方式(inner、outer、left outer…默认为inner)
需要注意的是,对于 两个DataFrame的连接依据 这个参数,若两个DF中这个列名相同(且表示的含义相同),则直接指定列名即可;若两个DF中依据列名不相同,则需要显示指定两个DF进行连接的列,比如:DF1(column1) === DF2(column2)。
② 将数据根据地区,商品名称分组,并统计每个分组的总点击数

分组函数groupBy其实与RDD非常类似,可以直接通过指定DataFrame中包含的列作为分组依据。
通过使用count()函数,会为当前DF新增一列count列,即:

③ 按找地区进行分组,对每个地区的点击次数降序排列,取出每个地区的点击次数前三

这一步就是对窗口函数的使用了。与sql文中使用略微有一些区别。
首先,若想要在DataFrameAPI中使用窗口函数rank() over(),则必须导入以下两个包:
import org.apache.spark.sql.expressions.Window
import org.apache.spark.sql.functions._
然后我们就可以对DataFrame进行使用了。我们在sql文中使用rank() over()时是会为原本的数据表新增一列组内排名列,因此在DataFrame中,我们同样需要通过withColumn函数新增一列rank列,作为组内排名。
withColumn函数的参数为:
-
-
- 列名
- 列表达式(基于什么操作得到该列)
-
rank列是通过窗口函数得来的,所以列表达式即窗口函数的使用:

不同于sql文中直接通过rank() over( partition by column order by column)的形式,在DataFrameAPI中需要在over中调用Window来使用partitionBy以及orderBy功能。
通过DataFrameAPI进行rank函数使用说明大致为:
//导入窗口函数两个关键sparkSQL包
import org.apache.spark.sql.expressions.Window
import org.apache.spark.sql.functions._
DF
//创建rank列作为分组排名列
.withColumn(
columnName, //指定列名
rank().over(
//通过Window调用partitionBy & orderBy方法
Window
.partitionBy(c1) //指定分组依据
.orderBy(c2) //指定组内排序依据
)
)
④ 完整的sparkSQL代码
package main.spark_sql.practice
import org.apache.spark.SparkConf
import org.apache.spark.sql.SparkSession
//通过纯DataFrameAPI实现需求
object CodeWithDataFrameAPI {
def main(args: Array[String]): Unit = {
val sparkConf = new SparkConf().setMaster("local[*]").setAppName("spark-sql-practice")
val spark = SparkSession.builder().config(sparkConf).getOrCreate()
import spark.implicits._
//点击表
val userDF = spark.read.format("jdbc")
.option("url", "jdbc:mysql://localhost:3306/spark_sql_practice") //要读取的MySQL数据库
.option("drive", "com.mysql.jdbc.driver") //指定jdbc驱动
.option("user", "root") //指定用户名
.option("password", "D200504193010") //指定密码
.option("dbtable", "user_visit_action") //指定要读取的数据表
.load
//产品表
val productDF = spark.read.format("jdbc")
.option("url", "jdbc:mysql://localhost:3306/spark_sql_practice") //要读取的MySQL数据库
.option("drive", "com.mysql.jdbc.driver") //指定jdbc驱动
.option("user", "root") //指定用户名
.option("password", "D200504193010") //指定密码
.option("dbtable", "product_info") //指定要读取的数据表
.load
//城市表
val cityDF = spark.read.format("jdbc")
.option("url", "jdbc:mysql://localhost:3306/spark_sql_practice") //要读取的MySQL数据库
.option("drive", "com.mysql.jdbc.driver") //指定jdbc驱动
.option("user", "root") //指定用户名
.option("password", "D200504193010") //指定密码
.option("dbtable", "city_info") //指定要读取的数据表
.load
//TODO 查询所有点击记录,与city_info表连接,得到每个城市所在地区,同时与product_info表连接,得到点击的产品名称
//将点击记录与city_info表连接,得到每个城市所在地区,同时与product_info表连接,得到点击的产品名称
//最终结果为后续需求使用的基础数据表
val basicDF = userDF
//选出是点击的数据
.filter($"click_product_id" > -1)
//通过city_id字段连接城市表
.join(cityDF, "city_id")
//通过product_id字段连接商品表,由于两个df列名不同,需要通过===去指定
.join(productDF, userDF("click_product_id") === productDF("product_id"))
//TODO 按地区 & 商品id分组,统计出每个商品在每个地区的总点击次数
val groupDF = basicDF
//按地区 & 商品id分组
.groupBy("area", "product_name")
//统计商品在每个地区的总点击次数
.count()
groupDF.show()
//TODO 每个地区按照点击次数降序排列,并且取出前三名
import org.apache.spark.sql.expressions.Window
import org.apache.spark.sql.functions._
val resultDF = groupDF
//使用窗口函数进行按地区的分组
.withColumn(
"rank",
rank().over(
Window.partitionBy("area").orderBy($"count".desc)
)
)
//过滤选出前三
.filter($"rank" <= 3)
//TODO 城市备注需要使用自定义UDAF函数
spark.close()
}
}
(2)UDAF代码实现(sql文中实现)

UDAF函数在sql文中使用会更加自然,因此我们此处演示在纯sql文的sparkSQL中使用UDAF函数进行统计。
a 将(1)中的完整sql文拆分到不同的TempView中
由于我们使用UDAF是在按照区域、商品聚合时使用,为了使得sql文不显得过分冗长以及嵌套过于混乱,我们将sql文拆开,为每一组查询都建立一张临时表。
package main.spark_sql.practice
import org.apache.spark.SparkConf
import org.apache.spark.sql.SparkSession
//基于sql文实现需求
object CodeWithSQLText {
def main(args: Array[String]): Unit = {
val sparkConf = new SparkConf().setMaster("local[*]").setAppName("spark-sql-practice")
val spark = SparkSession.builder().config(sparkConf).getOrCreate()
//点击表
val userDF = spark.read.format("jdbc")
.option("url", "jdbc:mysql://localhost:3306/spark_sql_practice") //要读取的MySQL数据库
.option("drive", "com.mysql.jdbc.driver") //指定jdbc驱动
.option("user", "root") //指定用户名
.option("password", "D200504193010") //指定密码
.option("dbtable", "user_visit_action") //指定要读取的数据表
.load
//产品表
val productDF = spark.read.format("jdbc")
.option("url", "jdbc:mysql://localhost:3306/spark_sql_practice") //要读取的MySQL数据库
.option("drive", "com.mysql.jdbc.driver") //指定jdbc驱动
.option("user", "root") //指定用户名
.option("password", "D200504193010") //指定密码
.option("dbtable", "product_info") //指定要读取的数据表
.load
//城市表
val cityDF = spark.read.format("jdbc")
.option("url", "jdbc:mysql://localhost:3306/spark_sql_practice") //要读取的MySQL数据库
.option("drive", "com.mysql.jdbc.driver") //指定jdbc驱动
.option("user", "root") //指定用户名
.option("password", "D200504193010") //指定密码
.option("dbtable", "city_info") //指定要读取的数据表
.load
//直接通过sql文查询
//首先创建三个临时表,用于sql文使用
userDF.createOrReplaceTempView("user_visit_action")
productDF.createOrReplaceTempView("product_info")
cityDF.createOrReplaceTempView("city_info")
//TODO 基础查询相关数据表
//连接后的表
spark.sql(
"""
|select
| u.*, p.product_name, c.city_name, c.area
| from
| user_visit_action u, product_info p, city_info c
| where
| u.city_id = c.city_id
| and u.click_product_id = p.product_id
| and u.click_product_id != -1""".stripMargin)
.createOrReplaceTempView("basicTable")
//按区域、商品进行分组
spark.sql(
"""
|select
| area, product_name, count(*) cnt
|from
| basicTable
|group by area, product_name
|""".stripMargin)
.createOrReplaceTempView("groupTable")
//按区域分组,每一组都对cnt降序排列,增加排名列字段
spark.sql(
"""
|select
| *,
| rank() over ( partition by area order by cnt desc ) rank_num
|from
| groupTable
|""".stripMargin)
.createOrReplaceTempView("orderTable")
//对每一组取出rank前三的数据
spark.sql(
"""
|select
| *
|from
| orderTable
|where
| rank_num <= 3
|""".stripMargin)
.createOrReplaceTempView("resultTable")
spark.sql("select * from resultTable").show()
//TODO 城市备注需要使用自定义UDAF函数
spark.close()
}
}
b 对分组sql进行修改,在按区域、商品分组时,通过UDAF统计城市备注信息

① UDAF的实现分析
我们首先分析一下需求:我们需要在按照区域、商品分组时,对传入的城市字段通过UDAF进行统计聚合。统计指标为:每个城市在当前区域、当前商品的分组中出现的占比。
因此,UDAF的流程大致为:
-
-
- 接收每个组的city_name字段作为IN
- 在缓冲区中,计算每个城市出现的次数count;同时计算所有城市出现的次数total
- 在缓冲区合并后,根据最后的结果,选出出现次数前两名的城市,计算其出现占比(count / total),然后计算其余城市出现占比
- 最后,将结果整合成String,作为OUT
-
以下是上述流程的一个简要图示:

② UDAF函数的实现(强类型UDAF)
UDAF的实现难点就在于IN、BUF、OUT参数如何设置,当我们搞清楚这些之后处理起来就非常轻松了。
/**
* UDAF需要使用的BUF缓冲区
*
* @param cityMap 用于统计每个分组中,每个城市分别出现的次数
* @param total 用于统计该分组中出现了多少城市
*/
case class Buf(var cityMap: mutable.Map[String, Int], var total: Int)
//统计城市备注使用的UDAF函数
class cityRemark extends Aggregator[String, Buf, String] {
//初始化缓冲区
override def zero: Buf = Buf(
mutable.Map.empty[String, Int], //创建一个空Map
0 //给城市总数设置初始值0
)
//缓冲区计算方法
//直接对buf属性进行更新!
override def reduce(buf: Buf, cityName: String): Buf = {
//根据输入的城市名作为key,更新buf中的城市Map,然后更新该区域总城市次数
buf.cityMap.update(cityName, buf.cityMap.getOrElse(cityName, 0) + 1)
buf.total += 1
buf
}
//合并两个缓冲区
override def merge(buf1: Buf, buf2: Buf): Buf = {
//合并到buf1中
val map1 = buf1.cityMap
val map2 = buf2.cityMap
//把buf2的cityMap合并到buf1的cityMap中,即更新buf1的map属性
map2.foreach(kv => map1.update(kv._1, map1.getOrElse(kv._1, 0) + kv._2))
//更新buf1的total属性
buf1.total += buf2.total
buf1
}
//缓冲区合并结束后,计算最终结果
//取出城市出现前二名,剩下的统一用其余表示,计算各自的百分比,整合成String输出
override def finish(resultBuf: Buf): String = {
//取出buf中城市的总出现次数
val total = resultBuf.total
//对buf内数据按value进行降序排列
val resultMap = resultBuf.cityMap
val sortCityList = resultMap
.toList.sortWith(_._2 > _._2) //转换成List进行排序,当前一个元素更大时才满足排序规则(降序)
.take(2) //取出前两名
//用于存放结果的对象
val resultList = new ListBuffer[String]
//判断是否有多余两个的城市,若没有,则显示前两名即可;若有,则需要添加其余城市描述
val hasMore = resultMap.size > 2 //是否有多余两个城市
var countSum = 0 //用于统计两个城市的出现总次数,以在有多余城市时通过该值进行其余城市百分比计算
//先将前两个城市信息添加
sortCityList.foreach(
{
case (cityName, count) => {
val percent = 100 * count / total
countSum += count
resultList.append(s"${cityName} ${percent}%")
}
}
)
//若有多余两个城市,则添加其余相关信息
if (hasMore)
resultList.append(s"其余 ${(total - countSum) * 100 / total}%")
resultList.mkString(",")
}
override def bufferEncoder: Encoder[Buf] = Encoders.product
override def outputEncoder: Encoder[String] = Encoders.STRING
}
- 缓冲区BUF的设置:

我们在 b-① 中已经分析了,缓冲区中要统计每个城市各自出现的次数 & 所有城市出现的总次数。
因此BUF需要两个参数即可:
cityMap:一个mutable.Map,用于统计每个城市各自出现的次数
total:用于统计当前分组所有城市出现的总次数
- 核心计算函数reduce & finish的展示:
- reduce:

对于reduce & merge这类对缓冲区的操作,再次强调,请尽量直接修改缓冲区属性,并且返回新的缓冲区,不要创建新的缓冲区对象,这样才能优化spark的性能!
- finish:

该部分主要考虑代码健壮性。因为若有两个以上的城市才需要添加 “其余“ 信息,所以我们使用一个ListBuffer[String]作为结果接收的集合,用于随时新增信息。
是否有其余集合的判断是根据buffer中的 “cityMap.size > 2“ 来决定的,不是通过取出前两名后的集合(这不废话吗,取出前两名了size还能大于2??)。
同时,由于 “其余“ 信息的计算与前两名城市的计算不在同一作用域,我们需要countSum这个参数来记录前两名城市的总出现次数,以在其余信息计算中使用。
最后,将结果集合拼接成String返回即可。整体而言没有什么太大难度。
③ 在sql文中使用强类型UDAF

通过functions.udaf将强类型转换为弱类型,然后在sql文中直接使用即可。这样我们就可以在分组的同时完成对city_name的聚合功能了。
c 完整代码 & 运行结果
- 完整代码:
package main.spark_sql.practice
import org.apache.spark.SparkConf
import org.apache.spark.sql.{Encoder, Encoders, SparkSession, functions}
import org.apache.spark.sql.expressions.Aggregator
import scala.collection.mutable
import scala.collection.mutable.ListBuffer
//基于sql文实现需求
object CodeWithSQLText {
def main(args: Array[String]): Unit = {
val sparkConf = new SparkConf().setMaster("local[*]").setAppName("spark-sql-practice")
val spark = SparkSession.builder().config(sparkConf).getOrCreate()
//点击表
val userDF = spark.read.format("jdbc")
.option("url", "jdbc:mysql://localhost:3306/spark_sql_practice") //要读取的MySQL数据库
.option("drive", "com.mysql.jdbc.driver") //指定jdbc驱动
.option("user", "root") //指定用户名
.option("password", "D200504193010") //指定密码
.option("dbtable", "user_visit_action") //指定要读取的数据表
.load
//产品表
val productDF = spark.read.format("jdbc")
.option("url", "jdbc:mysql://localhost:3306/spark_sql_practice") //要读取的MySQL数据库
.option("drive", "com.mysql.jdbc.driver") //指定jdbc驱动
.option("user", "root") //指定用户名
.option("password", "D200504193010") //指定密码
.option("dbtable", "product_info") //指定要读取的数据表
.load
//城市表
val cityDF = spark.read.format("jdbc")
.option("url", "jdbc:mysql://localhost:3306/spark_sql_practice") //要读取的MySQL数据库
.option("drive", "com.mysql.jdbc.driver") //指定jdbc驱动
.option("user", "root") //指定用户名
.option("password", "D200504193010") //指定密码
.option("dbtable", "city_info") //指定要读取的数据表
.load
//直接通过sql文查询
//首先创建三个临时表,用于sql文使用
userDF.createOrReplaceTempView("user_visit_action")
productDF.createOrReplaceTempView("product_info")
cityDF.createOrReplaceTempView("city_info")
//TODO 基础查询相关数据表
//连接后的表
spark.sql(
"""
|select
| u.*, p.product_name, c.city_name, c.area
| from
| user_visit_action u, product_info p, city_info c
| where
| u.city_id = c.city_id
| and u.click_product_id = p.product_id
| and u.click_product_id != -1""".stripMargin)
.createOrReplaceTempView("basicTable")
//TODO 城市备注需要使用自定义UDAF函数
//注册自定义UDAF函数
spark.udf.register("cityRemark", functions.udaf(new cityRemark))
//按区域、商品进行分组,同时对city_name参数通过udaf统计
spark.sql(
"""
|select
| area, product_name, count(*) cnt, cityRemark(city_name)
|from
| basicTable
|group by area, product_name
|""".stripMargin)
.createOrReplaceTempView("groupTable")
//按区域分组,每一组都对cnt降序排列,增加排名列字段
spark.sql(
"""
|select
| *,
| rank() over ( partition by area order by cnt desc ) rank_num
|from
| groupTable
|""".stripMargin)
.createOrReplaceTempView("orderTable")
//对每一组取出rank前三的数据
spark.sql(
"""
|select
| *
|from
| orderTable
|where
| rank_num <= 3
|""".stripMargin)
.createOrReplaceTempView("resultTable")
spark.sql("select * from resultTable").show()
spark.close()
}
/**
* UDAF需要使用的BUF缓冲区
*
* @param cityMap 用于统计每个分组中,每个城市分别出现的次数
* @param total 用于统计该分组中出现了多少城市
*/
case class Buf(var cityMap: mutable.Map[String, Int], var total: Int)
//统计城市备注使用的UDAF函数
class cityRemark extends Aggregator[String, Buf, String] {
//初始化缓冲区
override def zero: Buf = Buf(
mutable.Map.empty[String, Int], //创建一个空Map
0 //给城市总数设置初始值0
)
//缓冲区计算方法
//直接对buf属性进行更新!
override def reduce(buf: Buf, cityName: String): Buf = {
//根据输入的城市名作为key,更新buf中的城市Map,然后更新该区域总城市次数
buf.cityMap.update(cityName, buf.cityMap.getOrElse(cityName, 0) + 1)
buf.total += 1
buf
}
//合并两个缓冲区
override def merge(buf1: Buf, buf2: Buf): Buf = {
//合并到buf1中
val map1 = buf1.cityMap
val map2 = buf2.cityMap
//把buf2的cityMap合并到buf1的cityMap中,即更新buf1的map属性
map2.foreach(kv => map1.update(kv._1, map1.getOrElse(kv._1, 0) + kv._2))
//更新buf1的total属性
buf1.total += buf2.total
buf1
}
//缓冲区合并结束后,计算最终结果
//取出城市出现前二名,剩下的统一用其余表示,计算各自的百分比,整合成String输出
override def finish(resultBuf: Buf): String = {
//取出buf中城市的总出现次数
val total = resultBuf.total
//对buf内数据按value进行降序排列
val resultMap = resultBuf.cityMap
val sortCityList = resultMap
.toList.sortWith(_._2 > _._2) //转换成List进行排序,当前一个元素更大时才满足排序规则(降序)
.take(2) //取出前两名
//用于存放结果的对象
val resultList = new ListBuffer[String]
//判断是否有多余两个的城市,若没有,则显示前两名即可;若有,则需要添加其余城市描述
val hasMore = resultMap.size > 2 //是否有多余两个城市
var countSum = 0 //用于统计两个城市的出现总次数,以在有多余城市时通过该值进行其余城市百分比计算
//先将前两个城市信息添加
sortCityList.foreach(
{
case (cityName, count) => {
val percent = 100 * count / total
countSum += count
resultList.append(s"${cityName} ${percent}%")
}
}
)
//若有多余两个城市,则添加其余相关信息
if (hasMore)
resultList.append(s"其余 ${(total - countSum) * 100 / total}%")
resultList.mkString(",")
}
override def bufferEncoder: Encoder[Buf] = Encoders.product
override def outputEncoder: Encoder[String] = Encoders.STRING
}
}
- 运行结果:

更多推荐


所有评论(0)