Java反射实现通用DAO案例

wen java案例 2

本文目录导读:

Java反射实现通用DAO案例

  1. 通用DAO接口设计
  2. 实体类注解定义
  3. 实体类示例
  4. 反射工具类
  5. 通用DAO抽象类实现
  6. 具体DAO实现
  7. 使用示例
  8. Maven依赖
  9. 核心优势

我来为您提供一个基于Java反射的通用DAO实现案例,这个案例将展示如何使用反射技术实现通用的数据库CRUD操作。

通用DAO接口设计

import java.util.List;
import java.util.Map;
/**
 * 通用DAO接口
 */
public interface GenericDao<T> {
    /**
     * 插入记录
     */
    int insert(T entity) throws Exception;
    /**
     * 更新记录
     */
    int update(T entity) throws Exception;
    /**
     * 根据ID删除记录
     */
    int delete(Object id) throws Exception;
    /**
     * 根据ID查询记录
     */
    T findById(Object id) throws Exception;
    /**
     * 查询所有记录
     */
    List<T> findAll() throws Exception;
    /**
     * 根据条件查询
     */
    List<T> findByCondition(Map<String, Object> conditions) throws Exception;
    /**
     * 批量插入
     */
    int batchInsert(List<T> entities) throws Exception;
}

实体类注解定义

import java.lang.annotation.*;
/**
 * 表名注解
 */
@Target(ElementType.TYPE)
@Retention(RetentionPolicy.RUNTIME)
public @interface TableName {
    String value();
}
import java.lang.annotation.*;
/**
 * 字段注解(用于映射数据库列名)
 */
@Target(ElementType.FIELD)
@Retention(RetentionPolicy.RUNTIME)
public @interface Column {
    String value() default "";
    boolean primaryKey() default false;
}

实体类示例

/**
 * 用户实体类
 */
@TableName("t_user")
public class User {
    @Column(value = "id", primaryKey = true)
    private Long id;
    @Column("username")
    private String username;
    @Column("password")
    private String password;
    @Column("email")
    private String email;
    @Column("age")
    private Integer age;
    // 不需要映射的字段
    private String temp;
    // Getter和Setter方法
    public Long getId() {
        return id;
    }
    public void setId(Long id) {
        this.id = id;
    }
    public String getUsername() {
        return username;
    }
    public void setUsername(String username) {
        this.username = username;
    }
    public String getPassword() {
        return password;
    }
    public void setPassword(String password) {
        this.password = password;
    }
    public String getEmail() {
        return email;
    }
    public void setEmail(String email) {
        this.email = email;
    }
    public Integer getAge() {
        return age;
    }
    public void setAge(Integer age) {
        this.age = age;
    }
}

反射工具类

import java.lang.reflect.Field;
import java.util.ArrayList;
import java.util.HashMap;
import java.util.List;
import java.util.Map;
/**
 * 反射工具类
 */
public class ReflectUtil {
    /**
     * 获取表名
     */
    public static String getTableName(Class<?> clazz) {
        TableName annotation = clazz.getAnnotation(TableName.class);
        if (annotation != null) {
            return annotation.value();
        }
        // 如果没有注解,使用类名小写作为表名
        return clazz.getSimpleName().toLowerCase();
    }
    /**
     * 获取主键字段
     */
    public static Field getPrimaryKeyField(Class<?> clazz) {
        Field[] fields = clazz.getDeclaredFields();
        for (Field field : fields) {
            field.setAccessible(true);
            Column column = field.getAnnotation(Column.class);
            if (column != null && column.primaryKey()) {
                return field;
            }
        }
        return null;
    }
    /**
     * 获取主键列名
     */
    public static String getPrimaryKeyColumn(Class<?> clazz) {
        Field field = getPrimaryKeyField(clazz);
        if (field != null) {
            return getColumnName(field);
        }
        return "id";
    }
    /**
     * 获取字段对应的列名
     */
    public static String getColumnName(Field field) {
        Column column = field.getAnnotation(Column.class);
        if (column != null && !column.value().isEmpty()) {
            return column.value();
        }
        return field.getName();
    }
    /**
     * 获取所有映射的字段列表
     */
    public static List<Field> getMappedFields(Class<?> clazz) {
        List<Field> mappedFields = new ArrayList<>();
        Field[] fields = clazz.getDeclaredFields();
        for (Field field : fields) {
            if (field.isAnnotationPresent(Column.class)) {
                mappedFields.add(field);
            }
        }
        return mappedFields;
    }
    /**
     * 获取字段的值
     */
    public static Object getFieldValue(Object obj, Field field) throws IllegalAccessException {
        field.setAccessible(true);
        return field.get(obj);
    }
    /**
     * 设置字段的值
     */
    public static void setFieldValue(Object obj, Field field, Object value) throws IllegalAccessException {
        field.setAccessible(true);
        field.set(obj, value);
    }
    /**
     * 创建实体实例
     */
    public static <T> T newInstance(Class<T> clazz) throws Exception {
        return clazz.newInstance();
    }
    /**
     * 将实体对象转换为列名-值映射
     */
    public static Map<String, Object> entityToColumnMap(Object entity) throws Exception {
        Map<String, Object> columnMap = new HashMap<>();
        Class<?> clazz = entity.getClass();
        List<Field> fields = getMappedFields(clazz);
        for (Field field : fields) {
            String columnName = getColumnName(field);
            Object value = getFieldValue(entity, field);
            columnMap.put(columnName, value);
        }
        return columnMap;
    }
    /**
     * 将ResultSet转换为实体对象(需要注释掉,因为不直接使用数据库)
     * 这里简化实现,由具体DAO实现时处理
     */
    public static <T> T resultSetToEntity(Object resultSet, Class<T> clazz) {
        // 实际项目中在这里处理ResultSet到实体的转换
        try {
            T entity = clazz.newInstance();
            return entity;
        } catch (Exception e) {
            throw new RuntimeException(e);
        }
    }
}

通用DAO抽象类实现

import javax.sql.DataSource;
import java.sql.Connection;
import java.sql.PreparedStatement;
import java.sql.ResultSet;
import java.sql.Statement;
import java.util.ArrayList;
import java.util.List;
import java.util.Map;
/**
 * 通用DAO抽象类 - 基于JDBC实现
 */
public abstract class AbstractGenericDao<T> implements GenericDao<T> {
    protected DataSource dataSource;
    public AbstractGenericDao(DataSource dataSource) {
        this.dataSource = dataSource;
    }
    @Override
    public int insert(T entity) throws Exception {
        Class<?> clazz = entity.getClass();
        String tableName = ReflectUtil.getTableName(clazz);
        // 获取映射字段
        List<Field> fields = ReflectUtil.getMappedFields(clazz);
        // 构建SQL语句
        StringBuilder columns = new StringBuilder();
        StringBuilder values = new StringBuilder();
        for (Field field : fields) {
            if (columns.length() > 0) {
                columns.append(", ");
                values.append(", ");
            }
            columns.append(ReflectUtil.getColumnName(field));
            values.append("?");
        }
        String sql = String.format("INSERT INTO %s (%s) VALUES (%s)", 
                tableName, columns.toString(), values.toString());
        try (Connection conn = dataSource.getConnection();
             PreparedStatement ps = conn.prepareStatement(sql, Statement.RETURN_GENERATED_KEYS)) {
            // 设置参数
            int index = 1;
            for (Field field : fields) {
                Object value = ReflectUtil.getFieldValue(entity, field);
                ps.setObject(index++, value);
            }
            int rows = ps.executeUpdate();
            // 获取生成的主键
            ResultSet rs = ps.getGeneratedKeys();
            if (rs.next()) {
                Field pkField = ReflectUtil.getPrimaryKeyField(clazz);
                if (pkField != null) {
                    Object pkValue = rs.getObject(1);
                    ReflectUtil.setFieldValue(entity, pkField, pkValue);
                }
            }
            return rows;
        }
    }
    @Override
    public int update(T entity) throws Exception {
        Class<?> clazz = entity.getClass();
        String tableName = ReflectUtil.getTableName(clazz);
        // 获取映射字段
        List<Field> fields = ReflectUtil.getMappedFields(clazz);
        Field pkField = ReflectUtil.getPrimaryKeyField(clazz);
        if (pkField == null) {
            throw new IllegalArgumentException("实体类没有定义主键字段");
        }
        // 构建SQL语句
        StringBuilder setClause = new StringBuilder();
        for (Field field : fields) {
            if (field == pkField) continue;
            if (setClause.length() > 0) {
                setClause.append(", ");
            }
            setClause.append(ReflectUtil.getColumnName(field)).append(" = ?");
        }
        String pkColumn = ReflectUtil.getPrimaryKeyColumn(clazz);
        String sql = String.format("UPDATE %s SET %s WHERE %s = ?", 
                tableName, setClause.toString(), pkColumn);
        try (Connection conn = dataSource.getConnection();
             PreparedStatement ps = conn.prepareStatement(sql)) {
            // 设置参数
            int index = 1;
            for (Field field : fields) {
                if (field == pkField) continue;
                Object value = ReflectUtil.getFieldValue(entity, field);
                ps.setObject(index++, value);
            }
            // 设置主键参数
            Object pkValue = ReflectUtil.getFieldValue(entity, pkField);
            ps.setObject(index, pkValue);
            return ps.executeUpdate();
        }
    }
    @Override
    public int delete(Object id) throws Exception {
        Class<?> clazz = getEntityClass();
        String tableName = ReflectUtil.getTableName(clazz);
        String pkColumn = ReflectUtil.getPrimaryKeyColumn(clazz);
        String sql = String.format("DELETE FROM %s WHERE %s = ?", tableName, pkColumn);
        try (Connection conn = dataSource.getConnection();
             PreparedStatement ps = conn.prepareStatement(sql)) {
            ps.setObject(1, id);
            return ps.executeUpdate();
        }
    }
    @Override
    public T findById(Object id) throws Exception {
        Class<?> clazz = getEntityClass();
        String tableName = ReflectUtil.getTableName(clazz);
        String pkColumn = ReflectUtil.getPrimaryKeyColumn(clazz);
        String sql = String.format("SELECT * FROM %s WHERE %s = ?", tableName, pkColumn);
        try (Connection conn = dataSource.getConnection();
             PreparedStatement ps = conn.prepareStatement(sql)) {
            ps.setObject(1, id);
            try (ResultSet rs = ps.executeQuery()) {
                if (rs.next()) {
                    return mapResultSetToEntity(rs, clazz);
                }
            }
        }
        return null;
    }
    @Override
    public List<T> findAll() throws Exception {
        Class<?> clazz = getEntityClass();
        String tableName = ReflectUtil.getTableName(clazz);
        String sql = String.format("SELECT * FROM %s", tableName);
        List<T> results = new ArrayList<>();
        try (Connection conn = dataSource.getConnection();
             PreparedStatement ps = conn.prepareStatement(sql);
             ResultSet rs = ps.executeQuery()) {
            while (rs.next()) {
                results.add(mapResultSetToEntity(rs, clazz));
            }
        }
        return results;
    }
    @Override
    public List<T> findByCondition(Map<String, Object> conditions) throws Exception {
        Class<?> clazz = getEntityClass();
        String tableName = ReflectUtil.getTableName(clazz);
        StringBuilder sql = new StringBuilder("SELECT * FROM " + tableName);
        if (conditions != null && !conditions.isEmpty()) {
            sql.append(" WHERE ");
            List<Object> params = new ArrayList<>();
            int count = 0;
            for (Map.Entry<String, Object> entry : conditions.entrySet()) {
                if (count > 0) {
                    sql.append(" AND ");
                }
                sql.append(entry.getKey()).append(" = ?");
                params.add(entry.getValue());
                count++;
            }
            try (Connection conn = dataSource.getConnection();
                 PreparedStatement ps = conn.prepareStatement(sql.toString())) {
                for (int i = 0; i < params.size(); i++) {
                    ps.setObject(i + 1, params.get(i));
                }
                List<T> results = new ArrayList<>();
                try (ResultSet rs = ps.executeQuery()) {
                    while (rs.next()) {
                        results.add(mapResultSetToEntity(rs, clazz));
                    }
                }
                return results;
            }
        }
        return findAll();
    }
    @Override
    public int batchInsert(List<T> entities) throws Exception {
        int rows = 0;
        for (T entity : entities) {
            rows += insert(entity);
        }
        return rows;
    }
    /**
     * 获取实体类类型
     */
    protected abstract Class<T> getEntityClass();
    /**
     * 将ResultSet转换为实体对象
     */
    protected T mapResultSetToEntity(ResultSet rs, Class<?> clazz) throws Exception {
        T entity = (T) clazz.newInstance();
        List<Field> fields = ReflectUtil.getMappedFields(clazz);
        for (Field field : fields) {
            String columnName = ReflectUtil.getColumnName(field);
            Object value = rs.getObject(columnName);
            if (value != null) {
                ReflectUtil.setFieldValue(entity, field, value);
            }
        }
        return entity;
    }
}

具体DAO实现

/**
 * 用户DAO的具体实现
 */
public class UserDao extends AbstractGenericDao<User> {
    public UserDao(DataSource dataSource) {
        super(dataSource);
    }
    @Override
    protected Class<User> getEntityClass() {
        return User.class;
    }
    // 可以添加特定于User的业务方法
    public User findByUsername(String username) throws Exception {
        List<User> users = findByCondition(Collections.singletonMap("username", username));
        return users.isEmpty() ? null : users.get(0);
    }
    public List<User> findUsersByAgeRange(int minAge, int maxAge) throws Exception {
        // 自定义查询逻辑
        Class<User> clazz = getEntityClass();
        String tableName = ReflectUtil.getTableName(clazz);
        String sql = String.format("SELECT * FROM %s WHERE age BETWEEN ? AND ?", tableName);
        try (Connection conn = dataSource.getConnection();
             PreparedStatement ps = conn.prepareStatement(sql)) {
            ps.setInt(1, minAge);
            ps.setInt(2, maxAge);
            List<User> users = new ArrayList<>();
            try (ResultSet rs = ps.executeQuery()) {
                while (rs.next()) {
                    users.add(mapResultSetToEntity(rs, clazz));
                }
            }
            return users;
        }
    }
}

使用示例

/**
 * 通用DAO使用示例
 */
public class GenericDaoExample {
    public static void main(String[] args) throws Exception {
        // 创建数据源(示例使用H2数据库)
        BasicDataSource dataSource = new BasicDataSource();
        dataSource.setDriverClassName("org.h2.Driver");
        dataSource.setUrl("jdbc:h2:mem:testdb");
        dataSource.setUsername("sa");
        dataSource.setPassword("");
        // 创建表
        createTable(dataSource);
        // 创建DAO实例
        UserDao userDao = new UserDao(dataSource);
        // 插入用户
        User user1 = new User();
        user1.setUsername("张三");
        user1.setPassword("123456");
        user1.setEmail("zhangsan@example.com");
        user1.setAge(25);
        int inserted = userDao.insert(user1);
        System.out.println("插入用户1: " + (inserted > 0 ? "成功" : "失败"));
        System.out.println("生成的ID: " + user1.getId());
        // 批量插入
        List<User> users = new ArrayList<>();
        for (int i = 1; i <= 5; i++) {
            User user = new User();
            user.setUsername("用户" + i);
            user.setPassword("pass" + i);
            user.setEmail("user" + i + "@example.com");
            user.setAge(20 + i);
            users.add(user);
        }
        int batchInserted = userDao.batchInsert(users);
        System.out.println("批量插入: " + batchInserted + " 条记录");
        // 查询所有用户
        System.out.println("\n=== 所有用户 ===");
        List<User> allUsers = userDao.findAll();
        for (User user : allUsers) {
            System.out.println("ID: " + user.getId() + ", 用户名: " + user.getUsername() + ", 年龄: " + user.getAge());
        }
        // 根据ID查询
        System.out.println("\n=== 根据ID查询 ===");
        User foundUser = userDao.findById(1L);
        if (foundUser != null) {
            System.out.println("找到用户: " + foundUser.getUsername());
        }
        // 更新用户
        System.out.println("\n=== 更新用户 ===");
        foundUser.setEmail("newemail@example.com");
        foundUser.setAge(26);
        int updated = userDao.update(foundUser);
        System.out.println("更新结果: " + (updated > 0 ? "成功" : "失败"));
        // 条件查询
        System.out.println("\n=== 条件查询 ===");
        Map<String, Object> conditions = new HashMap<>();
        conditions.put("age", 23);
        List<User> filteredUsers = userDao.findByCondition(conditions);
        System.out.println("找到 " + filteredUsers.size() + " 个年龄为23的用户");
        // 删除用户
        System.out.println("\n=== 删除用户 ===");
        int deleted = userDao.delete(1L);
        System.out.println("删除结果: " + (deleted > 0 ? "成功" : "失败"));
        // 使用自定义方法
        System.out.println("\n=== 自定义查询 ===");
        List<User> ageRangeUsers = userDao.findUsersByAgeRange(20, 23);
        System.out.println("年龄在20-23之间的用户: " + ageRangeUsers.size());
    }
    private static void createTable(DataSource dataSource) throws Exception {
        String sql = "CREATE TABLE t_user (id BIGINT AUTO_INCREMENT PRIMARY KEY, " +
                     "username VARCHAR(50), password VARCHAR(100), email VARCHAR(100), age INT)";
        try (Connection conn = dataSource.getConnection();
             Statement stmt = conn.createStatement()) {
            stmt.execute(sql);
        }
    }
}

Maven依赖

<dependencies>
    <!-- H2数据库(示例使用) -->
    <dependency>
        <groupId>com.h2database</groupId>
        <artifactId>h2</artifactId>
        <version>2.1.214</version>
    </dependency>
    <!-- 连接池 -->
    <dependency>
        <groupId>commons-dbcp</groupId>
        <artifactId>commons-dbcp</artifactId>
        <version>1.4</version>
    </dependency>
</dependencies>

核心优势

  1. 减少重复代码:只需继承AbstractGenericDao即可获得CRUD功能
  2. 灵活性高:通过注解配置实体映射关系
  3. 类型安全:使用泛型确保类型安全
  4. 可扩展性:可以轻松添加其他通用操作
  5. 易于维护:修改数据库映射关系只需修改实体类注解

这个通用DAO框架通过Java反射技术,实现了对数据库操作的抽象和封装,大大简化了企业应用开发中的数据访问层代码。

抱歉,评论功能暂时关闭!