首页
学习
活动
专区
圈层
工具
发布
社区首页 >专栏 >MyBatis插件编写

MyBatis插件编写

作者头像
心平气和
发布2021-01-29 15:43:36
发布2021-01-29 15:43:36
9350
举报

一 、什么是MyBatis插件

Mybatis是一个操作数据库的工具,在一些场景下应用有些自定义的需求,在数据库整个执行流程上需有一些插入点可以接入自己的逻辑,如针对数据库敏感字段加密,分页等,因此MyBatis在设计的时候就采取发插件化的设计,可以让应用加入自己的逻辑。

今天我们来编写一个示例性的插件,这个插件的作用就是针对指定敏感字段入库时进行base64加密,出库时进行basex64解密,以保证数据库在脱库的情况下都不会发生泄漏,当然算法的安全性不是这篇文章的重点。

二、编写插件的大概步骤

今天的示例是在SpringBoot中编写,编写MyBatis的插件大概步骤如下:

1、实现Interceptor接口;

主要实现intercept和plugin方法

intercept方法的接口只有一个Invocation参数

代码语言:javascript
复制
public class Invocation {

  private Object target;
  private Method method;
  private Object[] args;

  public Invocation(Object target, Method method, Object[] args) {
    this.target = target;
    this.method = method;
    this.args = args;
  }

  public Object getTarget() {
    return target;
  }

  public Method getMethod() {
    return method;
  }

  public Object[] getArgs() {
    return args;
  }

  public Object proceed() throws InvocationTargetException, IllegalAccessException {
    return method.invoke(target, args);
  }

}

其中target表示当前拦截的对象实例,其它参数通过字面就知道含义了。

plugin按默认实现就可以了:

代码语言:javascript
复制
@Override
    public Object plugin(Object o) {
        return Plugin.wrap(o, this);
    }

2、为类加上Intercepts注解

注释格式如下:

代码语言:javascript
复制
@Intercepts({@Signature(type = ParameterHandler.class, method = "setParameters", args = {PreparedStatement.class}),
        @Signature(type = Executor.class, method = "query", args = {MappedStatement.class, Object.class, RowBounds.class, ResultHandler.class}),
        @Signature(type = ResultSetHandler.class, method = "handleResultSets", args = {Statement.class})})

这个注解包含的内容就是Signature注解数组,Signature注解包含3个参数:

type:作用在哪个类上,目前只支持以下4个类

Executor:执行Sql

ParameterHandler:参数处理

ResultSetHandler:结果处理

StatementHandler:对应JDBC中statement对像的操作

method:哪个方法,这个需要看下Mybatis相关代码,知道大概的执行流程

args:方法的签名,因为重载的原因,可能有多个同名方法,因此需要指定具体参数类型,方便MyBatis定位到具体的方法。

三、编写插件的具体操作步骤

1、编写注解的接口

代码语言:javascript
复制
@Retention(RetentionPolicy.RUNTIME)
@Documented
@Target({ElementType.FIELD})
@Inherited
public @interface Encryption {
}

这个注解标在属性上,如果标了我们才进行加解密。

2、编写插件代码

代码语言:javascript
复制
@Component
@Intercepts({@Signature(type = ParameterHandler.class, method = "setParameters", args = {PreparedStatement.class}),
        @Signature(type = Executor.class, method = "query", args = {MappedStatement.class, Object.class, RowBounds.class, ResultHandler.class}),
        @Signature(type = ResultSetHandler.class, method = "handleResultSets", args = {Statement.class})})
public class MyBatisDemoInterceptor implements Interceptor {
    @Override
    public Object intercept(Invocation invocation) throws Throwable {
        // 入参
        if (invocation.getTarget() instanceof ParameterHandler) {
            return encrypt(invocation);
        }

        // 出参
        if (invocation.getTarget() instanceof ResultSetHandler) {
            return decrypt(invocation);
        }

        return invocation.proceed();
    }

    @Override
    public Object plugin(Object o) {
        return Plugin.wrap(o, this);
    }

    @Override
    public void setProperties(Properties properties) {

    }

    /**
     * 加密
     * @param invocation
     * @return
     * @throws Exception
     */
    private Object encrypt(Invocation invocation) throws Exception{
        // 获取参数对象
        ParameterHandler parameterHandler = (ParameterHandler) invocation.getTarget();
        Field paramField = parameterHandler.getClass().getDeclaredField("parameterObject");
        if (null == paramField) {
            return invocation.proceed();
        }
        paramField.setAccessible(true);
        Object parameterObject = paramField.get(parameterHandler);
        if (null == parameterObject) {
            return invocation.proceed();
        }

        Class<?> clazz = parameterObject.getClass();
        Field[] fields = clazz.getDeclaredFields();
        for (Field field : fields) {
            field.setAccessible(true);

            Object value = field.get(parameterObject);
            if (null == value) {
                continue;
            }
            //获取字段上的注解
            Encryption encryption = field.getAnnotation(Encryption.class);
            if (null == encryption) {
                continue;
            }
            String text = Base64.getEncoder().encodeToString(((String)value).getBytes());
            field.set(parameterObject, text);
        }

        return invocation.proceed();
    }

    /**
     * 解密
     * @param invocation
     * @return
     * @throws Exception
     */
    private Object decrypt(Invocation invocation) throws Exception{
        Object result = invocation.proceed();
        if (null == result) {
            return result;
        }

        if (result instanceof List) {
            List batchResult = (List) result;
            if (CollectionUtils.isEmpty(batchResult)) {
                return result;
            }

            for (Object singleResult : batchResult) {
                decryptItem(singleResult);
            }
            return batchResult;
        }
        return invocation.proceed();
    }

    /**
     * 解密单个对象
     * @param object
     * @return
     */
    private void decryptItem(Object object){
        Class<?> clazz = object.getClass();
        Field[] fields = clazz.getDeclaredFields();
        for (Field field : fields) {
            field.setAccessible(true);

            try {
                Object value = field.get(object);
                if (null == value) {
                    continue;
                }

                Encryption encryption = field.getAnnotation(Encryption.class);
                if (null == encryption) {
                    continue;
                }
                byte[] decoded = Base64.getDecoder().decode(((String) value).getBytes());
                String text = new String(decoded);
                field.set(object, text);
            }catch (Exception ex){

            }
        }
    }
}

代码不是太复杂,我们在intercept方法中根据invocation.getTarget()判断是 ParameterHandler还是ResultSetHandler,如果是前者则加密,后者则解密;

加、解密的时候都是先遍历对象,看哪个字段有加Encryption注解,如果有则对其进行相应处理。

3、编写POJO,用于数据库中操作数据

代码语言:javascript
复制
public class PostPO {
    private int id;
    private String title;
    @Encryption
    private Object author;
    private Long createTime;
    private Long offset=0L;
    private Long limit=10L;

    public String getTitle() {
        return title;
    }

    public void setTitle(String title) {
        this.title = title;
    }

    public void setId(int id) {
        this.id = id;
    }

    public void setAuthor(Object author) {
        this.author = author;
    }

    public void setCreateTime(Long createTime) {
        this.createTime = createTime;
    }

    public void setOffset(Long offset) {
        this.offset = offset;
    }

    public void setLimit(Long limit) {
        this.limit = limit;
    }

    public int getId() {

        return id;
    }

    public Object getAuthor() {
        return author;
    }

    public Long getCreateTime() {
        return createTime;
    }

    public Long getOffset() {
        return offset;
    }

    public Long getLimit() {
        return limit;
    }
}

4、编写dao接口访问数据库

代码语言:javascript
复制
@Mapper
public interface PostDao {
    List<PostPO> list(PostPO postPO);
}

Xml如下:

代码语言:javascript
复制
<?xml version="1.0" encoding="UTF-8" ?>
<!DOCTYPE mapper PUBLIC "-//mybatis.org//DTD Mapper 3.0//EN" "http://mybatis.org/dtd/mybatis-3-mapper.dtd" >
<mapper namespace="com.springbootdemo.dao.PostDao">
    <resultMap id="BaseResultMap" type="Post">
        <id column="id" property="id" jdbcType="INTEGER"/>
        <result column="title" property="title" jdbcType="VARCHAR"/>
        <result column="author" property="author" jdbcType="VARCHAR"/>
        <result column="create_time" property="createTime" jdbcType="BIGINT"/>
    </resultMap>
    <sql id="Table_Name">
        post
    </sql>
    <select id="list" resultMap="BaseResultMap" parameterType="Post">
        select
        id,title,author,create_time
        from
        <include refid="Table_Name"/>
        <where>
            <if test="title !=null and title!=''">
                and title like CONCAT(#{title},'%')
            </if>
        </where>

        <if test="offset!=null and limit!=null and limit>0">
            limit #{offset},#{limit}
        </if>
    </select>
</mapper>

表结构如下:

代码语言:javascript
复制
CREATE TABLE `post` (
  `id` int(11) NOT NULL AUTO_INCREMENT,
  `title` varchar(255) DEFAULT NULL,
  `author` varchar(255) DEFAULT NULL,
  `create_time` bigint(255) DEFAULT NULL,
  PRIMARY KEY (`id`)
) ENGINE=InnoDB AUTO_INCREMENT=2 DEFAULT CHARSET=utf8;

然后就可以编写单元测试看下数据是否OK了。

本文参与 腾讯云自媒体同步曝光计划,分享自微信公众号。
原始发表:2021-01-26,如有侵权请联系 cloudcommunity@tencent.com 删除
问题归档专栏文章快讯文章归档关键词归档开发者手册归档开发者手册 Section 归档