feat(wm-iot): #30 设备影子服务+OTA固件升级完整实现

- entity: DeviceShadow, OtaFirmware, OtaTask, OtaUpgradeRecord
- service: DeviceShadowService (Redis Hash + TTL 24h, 上报/期望/delta/离线检测/批量查询)
- service: OtaService (固件管理/升级任务创建按批次/进度追踪/结果统计)
- controller: DeviceShadowController /api/iot/shadow/*
- controller: OtaController /api/iot/ota/*
- config: RedisConfig (RedisTemplate JSON序列化)
- SQL DDL: iot_ota_firmware, iot_ota_task, iot_ota_upgrade_record
- test: DeviceShadowServiceTest + OtaServiceTest (mock Redis + mock DB)
This commit is contained in:
2026-06-14 18:22:28 +08:00
parent 69fd9d7c41
commit 9763a5700d
13 changed files with 1858 additions and 29 deletions
+5
View File
@@ -13,5 +13,10 @@
<dependency><groupId>org.springframework.boot</groupId><artifactId>spring-boot-starter-data-redis</artifactId></dependency>
<dependency><groupId>org.postgresql</groupId><artifactId>postgresql</artifactId></dependency>
<dependency><groupId>net.postgis</groupId><artifactId>postgis-jdbc</artifactId></dependency>
<dependency>
<groupId>org.springframework.boot</groupId>
<artifactId>spring-boot-starter-test</artifactId>
<scope>test</scope>
</dependency>
</dependencies>
</project>
@@ -0,0 +1,59 @@
package com.water.iot.config;
import com.fasterxml.jackson.annotation.JsonAutoDetect;
import com.fasterxml.jackson.annotation.PropertyAccessor;
import com.fasterxml.jackson.databind.ObjectMapper;
import com.fasterxml.jackson.databind.SerializationFeature;
import com.fasterxml.jackson.datatype.jsr310.JavaTimeModule;
import org.springframework.context.annotation.Bean;
import org.springframework.context.annotation.Configuration;
import org.springframework.data.redis.connection.RedisConnectionFactory;
import org.springframework.data.redis.core.RedisTemplate;
import org.springframework.data.redis.core.StringRedisTemplate;
import org.springframework.data.redis.serializer.GenericJackson2JsonRedisSerializer;
import org.springframework.data.redis.serializer.StringRedisSerializer;
/**
* Redis配置
* - StringRedisTemplate: key/value均为String序列化(用于设备影子Hash存储)
* - RedisTemplate<String, Object>: key为String,value为JSON序列化(通用)
*/
@Configuration
public class RedisConfig {
@Bean
public StringRedisTemplate stringRedisTemplate(RedisConnectionFactory factory) {
StringRedisTemplate template = new StringRedisTemplate();
template.setConnectionFactory(factory);
return template;
}
@Bean
public RedisTemplate<String, Object> redisTemplate(RedisConnectionFactory factory) {
RedisTemplate<String, Object> template = new RedisTemplate<>();
template.setConnectionFactory(factory);
// JSON序列化器
ObjectMapper mapper = new ObjectMapper();
mapper.setVisibility(PropertyAccessor.ALL, JsonAutoDetect.Visibility.ANY);
mapper.activateDefaultTyping(
mapper.getPolymorphicTypeValidator(),
ObjectMapper.DefaultTyping.NON_FINAL
);
mapper.registerModule(new JavaTimeModule());
mapper.disable(SerializationFeature.WRITE_DATES_AS_TIMESTAMPS);
GenericJackson2JsonRedisSerializer jsonSerializer = new GenericJackson2JsonRedisSerializer(mapper);
StringRedisSerializer stringSerializer = new StringRedisSerializer();
// key序列化
template.setKeySerializer(stringSerializer);
template.setHashKeySerializer(stringSerializer);
// value序列化
template.setValueSerializer(jsonSerializer);
template.setHashValueSerializer(jsonSerializer);
template.afterPropertiesSet();
return template;
}
}
@@ -0,0 +1,107 @@
package com.water.iot.controller;
import com.water.common.core.result.R;
import com.water.iot.entity.DeviceShadow;
import com.water.iot.service.DeviceShadowService;
import io.swagger.v3.oas.annotations.Operation;
import io.swagger.v3.oas.annotations.tags.Tag;
import lombok.RequiredArgsConstructor;
import org.springframework.web.bind.annotation.*;
import java.util.List;
import java.util.Map;
@Tag(name = "设备影子")
@RestController
@RequestMapping("/api/iot/shadow")
@RequiredArgsConstructor
public class DeviceShadowController {
private final DeviceShadowService shadowService;
@Operation(summary = "获取设备影子")
@GetMapping("/{deviceSn}")
public R<DeviceShadow> getShadow(@PathVariable String deviceSn) {
DeviceShadow shadow = shadowService.getShadow(deviceSn);
if (shadow == null) {
return R.fail(404, "设备影子不存在: " + deviceSn);
}
return R.ok(shadow);
}
@Operation(summary = "更新设备上报状态")
@PostMapping("/{deviceSn}/reported")
public R<String> updateReported(@PathVariable String deviceSn,
@RequestBody Map<String, Object> state) {
shadowService.updateReported(deviceSn, state);
return R.ok("上报状态已更新");
}
@Operation(summary = "更新期望状态(云端下发)")
@PostMapping("/{deviceSn}/desired")
public R<String> updateDesired(@PathVariable String deviceSn,
@RequestBody Map<String, Object> desired) {
try {
com.fasterxml.jackson.databind.ObjectMapper mapper = new com.fasterxml.jackson.databind.ObjectMapper();
String json = mapper.writeValueAsString(desired);
shadowService.updateDesired(deviceSn, json);
return R.ok("期望状态已下发");
} catch (Exception e) {
return R.fail("序列化失败: " + e.getMessage());
}
}
@Operation(summary = "获取上报状态")
@GetMapping("/{deviceSn}/reported")
public R<String> getReportedState(@PathVariable String deviceSn) {
String state = shadowService.getReportedState(deviceSn);
return R.ok(state);
}
@Operation(summary = "获取期望状态")
@GetMapping("/{deviceSn}/desired")
public R<String> getDesiredState(@PathVariable String deviceSn) {
String state = shadowService.getDesiredState(deviceSn);
return R.ok(state);
}
@Operation(summary = "获取差异状态(delta)")
@GetMapping("/{deviceSn}/delta")
public R<String> getDeltaState(@PathVariable String deviceSn) {
String state = shadowService.getDeltaState(deviceSn);
return R.ok(state);
}
@Operation(summary = "批量查询设备影子")
@PostMapping("/batch")
public R<List<DeviceShadow>> batchGetShadows(@RequestBody List<String> deviceSns) {
return R.ok(shadowService.batchGetShadows(deviceSns));
}
@Operation(summary = "检查设备在线状态")
@GetMapping("/{deviceSn}/online")
public R<Boolean> checkOnline(@PathVariable String deviceSn,
@RequestParam(defaultValue = "30") int thresholdMinutes) {
return R.ok(shadowService.checkOnline(deviceSn, thresholdMinutes));
}
@Operation(summary = "批量检测在线状态")
@PostMapping("/online/batch")
public R<Map<String, Boolean>> batchCheckOnline(@RequestBody List<String> deviceSns,
@RequestParam(defaultValue = "30") int thresholdMinutes) {
return R.ok(shadowService.batchCheckOnline(deviceSns, thresholdMinutes));
}
@Operation(summary = "在线设备数量")
@GetMapping("/online/count")
public R<Long> countOnline() {
return R.ok(shadowService.countOnlineDevices());
}
@Operation(summary = "删除设备影子")
@DeleteMapping("/{deviceSn}")
public R<String> deleteShadow(@PathVariable String deviceSn) {
shadowService.deleteShadow(deviceSn);
return R.ok("影子已删除");
}
}
@@ -0,0 +1,179 @@
package com.water.iot.controller;
import com.water.common.core.result.R;
import com.water.iot.entity.OtaFirmware;
import com.water.iot.entity.OtaTask;
import com.water.iot.entity.OtaUpgradeRecord;
import com.water.iot.service.OtaService;
import io.swagger.v3.oas.annotations.Operation;
import io.swagger.v3.oas.annotations.tags.Tag;
import lombok.RequiredArgsConstructor;
import org.springframework.web.bind.annotation.*;
import java.util.List;
import java.util.Map;
@Tag(name = "OTA固件升级")
@RestController
@RequestMapping("/api/iot/ota")
@RequiredArgsConstructor
public class OtaController {
private final OtaService otaService;
// ========== 固件管理 ==========
@Operation(summary = "创建固件版本")
@PostMapping("/firmware")
public R<OtaFirmware> createFirmware(@RequestBody OtaFirmware firmware) {
return R.ok(otaService.createFirmware(firmware));
}
@Operation(summary = "发布固件")
@PostMapping("/firmware/{id}/publish")
public R<String> publishFirmware(@PathVariable Long id,
@RequestParam String publishedBy) {
otaService.publishFirmware(id, publishedBy);
return R.ok("固件已发布");
}
@Operation(summary = "废弃固件")
@PostMapping("/firmware/{id}/deprecate")
public R<String> deprecateFirmware(@PathVariable Long id) {
otaService.deprecateFirmware(id);
return R.ok("固件已废弃");
}
@Operation(summary = "查询固件详情")
@GetMapping("/firmware/{id}")
public R<OtaFirmware> getFirmware(@PathVariable Long id) {
OtaFirmware firmware = otaService.getFirmware(id);
if (firmware == null) {
return R.fail(404, "固件不存在: " + id);
}
return R.ok(firmware);
}
@Operation(summary = "按模型查询固件列表")
@GetMapping("/firmware/model/{modelId}")
public R<List<OtaFirmware>> listByModel(@PathVariable Long modelId) {
return R.ok(otaService.listFirmwareByModel(modelId));
}
@Operation(summary = "查询所有固件(分页)")
@GetMapping("/firmware")
public R<List<OtaFirmware>> listAllFirmware(@RequestParam(defaultValue = "1") int page,
@RequestParam(defaultValue = "10") int size) {
return R.ok(otaService.listAllFirmware(page, size));
}
@Operation(summary = "设备查询可用固件")
@GetMapping("/firmware/available/{deviceSn}")
public R<OtaFirmware> getAvailableFirmware(@PathVariable String deviceSn) {
OtaFirmware firmware = otaService.getAvailableFirmware(deviceSn);
return R.ok(firmware);
}
// ========== 升级任务管理 ==========
@Operation(summary = "创建升级任务(指定设备列表)")
@PostMapping("/task")
public R<OtaTask> createUpgradeTask(@RequestBody Map<String, Object> body) {
Long firmwareId = ((Number) body.get("firmwareId")).longValue();
@SuppressWarnings("unchecked")
List<Number> deviceIdNums = (List<Number>) body.get("deviceIds");
List<Long> deviceIds = deviceIdNums.stream().map(Number::longValue).toList();
int batchSize = body.containsKey("batchSize") ? ((Number) body.get("batchSize")).intValue() : 10;
String createdBy = (String) body.getOrDefault("createdBy", "system");
return R.ok(otaService.createUpgradeTask(firmwareId, deviceIds, batchSize, createdBy));
}
@Operation(summary = "按条件创建升级任务")
@PostMapping("/task/filter")
public R<OtaTask> createTaskByFilter(@RequestBody Map<String, Object> body) {
Long firmwareId = ((Number) body.get("firmwareId")).longValue();
String deviceType = (String) body.get("deviceType");
String area = (String) body.get("area");
int batchSize = body.containsKey("batchSize") ? ((Number) body.get("batchSize")).intValue() : 10;
String createdBy = (String) body.getOrDefault("createdBy", "system");
return R.ok(otaService.createUpgradeTaskByFilter(firmwareId, deviceType, area, batchSize, createdBy));
}
@Operation(summary = "启动升级任务")
@PostMapping("/task/{id}/start")
public R<String> startTask(@PathVariable Long id) {
otaService.startTask(id);
return R.ok("任务已启动");
}
@Operation(summary = "取消升级任务")
@PostMapping("/task/{id}/cancel")
public R<String> cancelTask(@PathVariable Long id) {
otaService.cancelTask(id);
return R.ok("任务已取消");
}
@Operation(summary = "查询升级任务详情")
@GetMapping("/task/{id}")
public R<OtaTask> getTask(@PathVariable Long id) {
OtaTask task = otaService.getTask(id);
if (task == null) {
return R.fail(404, "任务不存在: " + id);
}
return R.ok(task);
}
@Operation(summary = "查询升级任务列表")
@GetMapping("/task")
public R<List<OtaTask>> listTasks(@RequestParam(defaultValue = "1") int page,
@RequestParam(defaultValue = "10") int size) {
return R.ok(otaService.listTasks(page, size));
}
// ========== 进度追踪 ==========
@Operation(summary = "更新升级进度")
@PostMapping("/progress/{recordId}")
public R<String> updateProgress(@PathVariable Long recordId,
@RequestParam int progress) {
otaService.updateProgress(recordId, progress);
return R.ok("进度已更新");
}
@Operation(summary = "标记升级成功")
@PostMapping("/record/{recordId}/success")
public R<String> markSuccess(@PathVariable Long recordId,
@RequestParam Long deviceId,
@RequestParam String toVersion) {
otaService.markSuccess(recordId, deviceId, toVersion);
return R.ok("已标记成功");
}
@Operation(summary = "标记升级失败")
@PostMapping("/record/{recordId}/fail")
public R<String> markFailed(@PathVariable Long recordId,
@RequestParam String reason) {
otaService.markFailed(recordId, reason);
return R.ok("已标记失败");
}
@Operation(summary = "查询任务统计")
@GetMapping("/task/{id}/statistics")
public R<Map<String, Object>> getTaskStatistics(@PathVariable Long id) {
return R.ok(otaService.getTaskStatistics(id));
}
@Operation(summary = "查询任务升级记录")
@GetMapping("/task/{id}/records")
public R<List<OtaUpgradeRecord>> getTaskRecords(@PathVariable Long id) {
return R.ok(otaService.getTaskRecords(id));
}
@Operation(summary = "查询设备升级历史")
@GetMapping("/device/{deviceSn}/history")
public R<List<OtaUpgradeRecord>> getDeviceHistory(@PathVariable String deviceSn) {
return R.ok(otaService.getDeviceUpgradeHistory(deviceSn));
}
}
@@ -0,0 +1,40 @@
package com.water.iot.entity;
import lombok.Data;
import java.time.LocalDateTime;
/**
* 设备影子实体
* Redis Hash存储:key=shadow:{deviceSn}
* fields: reported(JSON), desired(JSON), lastReportTime, online
*/
@Data
public class DeviceShadow {
/** 设备SN(主键) */
private String deviceSn;
/** 设备ID */
private Long deviceId;
/** 设备名称 */
private String deviceName;
/** 上报状态JSON */
private String reportedState;
/** 期望状态JSON(云端下发) */
private String desiredState;
/** 差异状态JSON */
private String deltaState;
/** 最后上报时间 */
private LocalDateTime lastReportTime;
/** 在线状态 */
private Boolean online;
/** 影子版本号 */
private Long version;
}
@@ -0,0 +1,44 @@
package com.water.iot.entity;
import lombok.Data;
import java.time.LocalDateTime;
/**
* OTA固件版本实体
*/
@Data
public class OtaFirmware {
/** 固件ID */
private Long id;
/** 设备模型ID */
private Long modelId;
/** 固件版本号 */
private String firmwareVersion;
/** 固件文件URL */
private String fileUrl;
/** 固件描述 */
private String description;
/** 状态: draft/published/deprecated */
private String status;
/** 文件MD5 */
private String md5;
/** 文件大小(bytes) */
private Long fileSize;
/** 发布人 */
private String publishedBy;
/** 发布时间 */
private LocalDateTime publishedAt;
private LocalDateTime createdAt;
private LocalDateTime updatedAt;
}
@@ -0,0 +1,59 @@
package com.water.iot.entity;
import lombok.Data;
import java.time.LocalDateTime;
/**
* OTA升级任务实体
*/
@Data
public class OtaTask {
/** 任务ID */
private Long id;
/** 固件ID */
private Long firmwareId;
/** 固件版本号 */
private String firmwareVersion;
/** 目标设备ID列表(JSON) */
private String targetDeviceIds;
/** 任务状态: pending/executing/completed/cancelled */
private String taskStatus;
/** 每批大小 */
private Integer batchSize;
/** 目标设备总数 */
private Integer totalDevices;
/** 成功数 */
private Integer successCount;
/** 失败数 */
private Integer failedCount;
/** 执行中数量 */
private Integer executingCount;
/** 目标设备类型(可选) */
private String targetType;
/** 目标区域(可选) */
private String targetArea;
/** 创建人 */
private String createdBy;
/** 执行开始时间 */
private LocalDateTime executedAt;
/** 完成时间 */
private LocalDateTime completedAt;
private LocalDateTime createdAt;
private LocalDateTime updatedAt;
}
@@ -0,0 +1,46 @@
package com.water.iot.entity;
import lombok.Data;
import java.time.LocalDateTime;
/**
* OTA升级记录实体(每台设备一条)
*/
@Data
public class OtaUpgradeRecord {
/** 记录ID */
private Long id;
/** 任务ID */
private Long taskId;
/** 设备ID */
private Long deviceId;
/** 设备SN */
private String deviceSn;
/** 升级前版本 */
private String fromVersion;
/** 升级目标版本 */
private String toVersion;
/** 状态: pending/executing/success/failed */
private String status;
/** 进度(0-100) */
private Integer progress;
/** 失败原因 */
private String failReason;
/** 开始时间 */
private LocalDateTime startedAt;
/** 完成时间 */
private LocalDateTime completedAt;
private LocalDateTime createdAt;
}
@@ -2,15 +2,25 @@ package com.water.iot.service;
import com.fasterxml.jackson.core.JsonProcessingException;
import com.fasterxml.jackson.databind.ObjectMapper;
import com.water.iot.entity.DeviceShadow;
import lombok.RequiredArgsConstructor;
import lombok.extern.slf4j.Slf4j;
import org.springframework.data.redis.core.StringRedisTemplate;
import org.springframework.jdbc.core.BeanPropertyRowMapper;
import org.springframework.jdbc.core.JdbcTemplate;
import org.springframework.stereotype.Service;
import java.util.Map;
import java.time.LocalDateTime;
import java.time.temporal.ChronoUnit;
import java.util.*;
import java.util.concurrent.TimeUnit;
import java.util.stream.Collectors;
/**
* 设备影子服务
* Redis Hash存储:key=shadow:{deviceSn},TTL 24h
* fields: reported(JSON), desired(JSON), lastReportTime, online, version
*/
@Slf4j
@Service
@RequiredArgsConstructor
@@ -18,32 +28,251 @@ public class DeviceShadowService {
private final StringRedisTemplate redisTemplate;
private final JdbcTemplate jdbcTemplate;
private final ObjectMapper mapper = new ObjectMapper();
private final ObjectMapper objectMapper;
private static final String SHADOW_PREFIX = "iot:shadow:";
private static final String SHADOW_PREFIX = "shadow:";
private static final long SHADOW_TTL_HOURS = 24;
/** 更新设备上报状态 */
/**
* 更新设备上报状态
* @param deviceSn 设备SN
* @param state 上报的状态数据(Map)
*/
public void updateReported(String deviceSn, Map<String, Object> state) {
try {
String key = SHADOW_PREFIX + deviceSn;
String json = mapper.writeValueAsString(state);
redisTemplate.opsForHash().put(key, "reported", json);
String json = objectMapper.writeValueAsString(state);
Map<String, String> fields = new HashMap<>();
fields.put("reported", json);
fields.put("lastReportTime", LocalDateTime.now().toString());
fields.put("online", "true");
// 递增版本号
Long version = redisTemplate.opsForHash().increment(key, "version", 1);
fields.put("version", String.valueOf(version));
redisTemplate.opsForHash().putAll(key, fields);
redisTemplate.expire(key, SHADOW_TTL_HOURS, TimeUnit.HOURS);
// 同步更新数据库设备最后上报时间
jdbcTemplate.update("UPDATE iot_device SET last_report_time = NOW() WHERE device_sn = ?", deviceSn);
// 计算delta(期望与上报的差异)
computeDelta(deviceSn);
// 同步更新数据库设备最后上报时间和在线状态
jdbcTemplate.update(
"UPDATE iot_device SET last_report_time = NOW(), status = 'online' WHERE device_sn = ?",
deviceSn
);
log.debug("Shadow reported updated: deviceSn={}, version={}", deviceSn, version);
} catch (JsonProcessingException e) {
log.error("Shadow update error: {}", e.getMessage());
log.error("Shadow update error for device {}: {}", deviceSn, e.getMessage());
}
}
/** 获取设备影子 */
public Map<Object, Object> getShadow(String deviceSn) {
return redisTemplate.opsForHash().entries(SHADOW_PREFIX + deviceSn);
/**
* 更新期望状态(云端→设备)
* @param deviceSn 设备SN
* @param desiredJson 期望状态JSON字符串
*/
public void updateDesired(String deviceSn, String desiredJson) {
String key = SHADOW_PREFIX + deviceSn;
Map<String, String> fields = new HashMap<>();
fields.put("desired", desiredJson);
Long version = redisTemplate.opsForHash().increment(key, "version", 1);
fields.put("version", String.valueOf(version));
redisTemplate.opsForHash().putAll(key, fields);
redisTemplate.expire(key, SHADOW_TTL_HOURS, TimeUnit.HOURS);
computeDelta(deviceSn);
log.debug("Shadow desired updated: deviceSn={}, version={}", deviceSn, version);
}
/** 更新期望状态(云端→设备) */
public void updateDesired(String deviceSn, String desiredJson) {
redisTemplate.opsForHash().put(SHADOW_PREFIX + deviceSn, "desired", desiredJson);
/**
* 获取设备影子完整数据
*/
public DeviceShadow getShadow(String deviceSn) {
String key = SHADOW_PREFIX + deviceSn;
Map<Object, Object> entries = redisTemplate.opsForHash().entries(key);
if (entries.isEmpty()) {
return null;
}
return mapToShadow(deviceSn, entries);
}
/**
* 获取设备上报状态
*/
public String getReportedState(String deviceSn) {
Object reported = redisTemplate.opsForHash().get(SHADOW_PREFIX + deviceSn, "reported");
return reported != null ? reported.toString() : null;
}
/**
* 获取设备期望状态
*/
public String getDesiredState(String deviceSn) {
Object desired = redisTemplate.opsForHash().get(SHADOW_PREFIX + deviceSn, "desired");
return desired != null ? desired.toString() : null;
}
/**
* 获取差异状态(delta = desired - reported)
*/
public String getDeltaState(String deviceSn) {
Object delta = redisTemplate.opsForHash().get(SHADOW_PREFIX + deviceSn, "delta");
return delta != null ? delta.toString() : null;
}
/**
* 批量查询设备影子
* @param deviceSns 设备SN列表
*/
public List<DeviceShadow> batchGetShadows(List<String> deviceSns) {
if (deviceSns == null || deviceSns.isEmpty()) {
return Collections.emptyList();
}
return deviceSns.stream()
.map(this::getShadow)
.filter(Objects::nonNull)
.collect(Collectors.toList());
}
/**
* 离线检测:检查指定设备是否在线(基于Redis TTL)
* @param deviceSn 设备SN
* @param offlineThresholdMinutes 离线阈值(分钟),超过该时间无上报视为离线
* @return true=在线, false=离线
*/
public boolean checkOnline(String deviceSn, int offlineThresholdMinutes) {
String key = SHADOW_PREFIX + deviceSn;
Object lastReportTimeStr = redisTemplate.opsForHash().get(key, "lastReportTime");
if (lastReportTimeStr == null) {
// 无上报记录,查DB
List<LocalDateTime> dbTimes = jdbcTemplate.queryForList(
"SELECT last_report_time FROM iot_device WHERE device_sn = ?",
LocalDateTime.class, deviceSn
);
if (dbTimes.isEmpty() || dbTimes.get(0) == null) return false;
lastReportTimeStr = dbTimes.get(0).toString();
}
LocalDateTime lastReport = LocalDateTime.parse(lastReportTimeStr.toString());
long minutesSinceReport = ChronoUnit.MINUTES.between(lastReport, LocalDateTime.now());
boolean isOnline = minutesSinceReport <= offlineThresholdMinutes;
// 更新Redis在线状态
redisTemplate.opsForHash().put(key, "online", String.valueOf(isOnline));
// 同步到DB
String dbStatus = isOnline ? "online" : "offline";
jdbcTemplate.update("UPDATE iot_device SET status = ? WHERE device_sn = ?", dbStatus, deviceSn);
return isOnline;
}
/**
* 批量离线检测(使用默认阈值30分钟)
*/
public Map<String, Boolean> batchCheckOnline(List<String> deviceSns) {
return batchCheckOnline(deviceSns, 30);
}
/**
* 批量离线检测
*/
public Map<String, Boolean> batchCheckOnline(List<String> deviceSns, int offlineThresholdMinutes) {
Map<String, Boolean> result = new LinkedHashMap<>();
for (String sn : deviceSns) {
result.put(sn, checkOnline(sn, offlineThresholdMinutes));
}
return result;
}
/**
* 删除设备影子(设备注销时)
*/
public void deleteShadow(String deviceSn) {
redisTemplate.delete(SHADOW_PREFIX + deviceSn);
}
/**
* 查询在线设备数量(Redis中有记录且online=true)
*/
public long countOnlineDevices() {
// 使用DB统计更可靠,这里提供Redis快速统计
List<Map<String, Object>> result = jdbcTemplate.queryForList(
"SELECT COUNT(*) as cnt FROM iot_device WHERE status = 'online'"
);
return result.isEmpty() ? 0 : ((Number) result.get(0).get("cnt")).longValue();
}
// ===== 私有方法 =====
/**
* 计算delta状态(期望与上报的差异)
*/
@SuppressWarnings("unchecked")
private void computeDelta(String deviceSn) {
String key = SHADOW_PREFIX + deviceSn;
try {
Object reportedObj = redisTemplate.opsForHash().get(key, "reported");
Object desiredObj = redisTemplate.opsForHash().get(key, "desired");
if (desiredObj == null) {
// 无期望状态,delta为空
redisTemplate.opsForHash().put(key, "delta", "{}");
return;
}
Map<String, Object> reported = reportedObj != null ?
objectMapper.readValue(reportedObj.toString(), Map.class) : Collections.emptyMap();
Map<String, Object> desired = objectMapper.readValue(desiredObj.toString(), Map.class);
// delta = desired中不同于reported的字段
Map<String, Object> delta = new LinkedHashMap<>();
for (Map.Entry<String, Object> entry : desired.entrySet()) {
Object reportedVal = reported.get(entry.getKey());
if (!Objects.equals(entry.getValue(), reportedVal)) {
delta.put(entry.getKey(), entry.getValue());
}
}
String deltaJson = objectMapper.writeValueAsString(delta);
redisTemplate.opsForHash().put(key, "delta", deltaJson);
} catch (Exception e) {
log.warn("Delta compute error for device {}: {}", deviceSn, e.getMessage());
}
}
/**
* Redis Hash → DeviceShadow
*/
private DeviceShadow mapToShadow(String deviceSn, Map<Object, Object> entries) {
DeviceShadow shadow = new DeviceShadow();
shadow.setDeviceSn(deviceSn);
shadow.setReportedState(entries.get("reported") != null ? entries.get("reported").toString() : null);
shadow.setDesiredState(entries.get("desired") != null ? entries.get("desired").toString() : null);
shadow.setDeltaState(entries.get("delta") != null ? entries.get("delta").toString() : null);
shadow.setOnline(entries.get("online") != null && "true".equals(entries.get("online").toString()));
shadow.setVersion(entries.get("version") != null ? Long.parseLong(entries.get("version").toString()) : 0L);
if (entries.get("lastReportTime") != null) {
shadow.setLastReportTime(LocalDateTime.parse(entries.get("lastReportTime").toString()));
}
// 从DB补充设备信息
List<Map<String, Object>> deviceInfo = jdbcTemplate.queryForList(
"SELECT id, device_name FROM iot_device WHERE device_sn = ?", deviceSn
);
if (!deviceInfo.isEmpty()) {
shadow.setDeviceId(((Number) deviceInfo.get(0).get("id")).longValue());
shadow.setDeviceName((String) deviceInfo.get(0).get("device_name"));
}
return shadow;
}
}
@@ -1,12 +1,30 @@
package com.water.iot.service;
import com.fasterxml.jackson.core.JsonProcessingException;
import com.fasterxml.jackson.core.type.TypeReference;
import com.fasterxml.jackson.databind.ObjectMapper;
import com.water.iot.entity.OtaFirmware;
import com.water.iot.entity.OtaTask;
import com.water.iot.entity.OtaUpgradeRecord;
import lombok.RequiredArgsConstructor;
import lombok.extern.slf4j.Slf4j;
import org.springframework.jdbc.core.BeanPropertyRowMapper;
import org.springframework.jdbc.core.JdbcTemplate;
import org.springframework.jdbc.support.GeneratedKeyHolder;
import org.springframework.jdbc.support.KeyHolder;
import org.springframework.stereotype.Service;
import org.springframework.transaction.annotation.Transactional;
import java.util.Map;
import java.sql.PreparedStatement;
import java.sql.Statement;
import java.time.LocalDateTime;
import java.util.*;
import java.util.stream.Collectors;
/**
* OTA固件升级服务
* 支持固件上传/版本管理/升级任务创建(按批次)/进度追踪/结果统计/设备查询可用固件
*/
@Slf4j
@Service
@RequiredArgsConstructor
@@ -14,21 +32,475 @@ public class OtaService {
private final JdbcTemplate jdbcTemplate;
private final DeviceShadowService shadowService;
private final ObjectMapper objectMapper;
/** 创建 OTA 升级任务 */
public void createUpgrade(Long modelId, String firmwareVersion, String firmwareUrl, String checkMd5) {
jdbcTemplate.update(
"INSERT INTO iot_device_event (device_id, device_sn, event_type, event_data) " +
"SELECT id, device_sn, 'ota', json_build_object('version',?, 'url',?, 'md5',?) " +
"FROM iot_device WHERE model_id = ? AND status = 'online'",
firmwareVersion, firmwareUrl, checkMd5, modelId);
log.info("OTA task created for model {}: version={}", modelId, firmwareVersion);
// ========== 固件管理 ==========
/**
* 创建固件版本
*/
@Transactional
public OtaFirmware createFirmware(OtaFirmware firmware) {
firmware.setStatus("draft");
firmware.setCreatedAt(LocalDateTime.now());
firmware.setUpdatedAt(LocalDateTime.now());
KeyHolder keyHolder = new GeneratedKeyHolder();
jdbcTemplate.update(connection -> {
var ps = connection.prepareStatement(
"INSERT INTO iot_ota_firmware (model_id, firmware_version, file_url, description, status, md5, file_size, created_at, updated_at) " +
"VALUES (?, ?, ?, ?, ?, ?, ?, NOW(), NOW())",
Statement.RETURN_GENERATED_KEYS
);
ps.setLong(1, firmware.getModelId());
ps.setString(2, firmware.getFirmwareVersion());
ps.setString(3, firmware.getFileUrl());
ps.setString(4, firmware.getDescription());
ps.setString(5, firmware.getStatus());
ps.setString(6, firmware.getMd5());
ps.setLong(7, firmware.getFileSize() != null ? firmware.getFileSize() : 0);
return ps;
}, keyHolder);
firmware.setId(keyHolder.getKey().longValue());
log.info("Firmware created: id={}, version={}", firmware.getId(), firmware.getFirmwareVersion());
return firmware;
}
/** 设备查询是否有待升级固件 */
public Map<String, Object> checkUpgrade(String deviceSn, String currentVersion) {
return jdbcTemplate.queryForMap(
"SELECT * FROM iot_device_event WHERE device_sn = ? AND event_type = 'ota' ORDER BY created_at DESC LIMIT 1",
deviceSn);
/**
* 发布固件
*/
@Transactional
public void publishFirmware(Long firmwareId, String publishedBy) {
int updated = jdbcTemplate.update(
"UPDATE iot_ota_firmware SET status = 'published', published_by = ?, published_at = NOW(), updated_at = NOW() WHERE id = ?",
publishedBy, firmwareId
);
if (updated == 0) throw new RuntimeException("固件不存在: " + firmwareId);
log.info("Firmware {} published by {}", firmwareId, publishedBy);
}
/**
* 废弃固件
*/
@Transactional
public void deprecateFirmware(Long firmwareId) {
jdbcTemplate.update(
"UPDATE iot_ota_firmware SET status = 'deprecated', updated_at = NOW() WHERE id = ?",
firmwareId
);
}
/**
* 查询固件详情
*/
public OtaFirmware getFirmware(Long firmwareId) {
List<OtaFirmware> list = jdbcTemplate.query(
"SELECT * FROM iot_ota_firmware WHERE id = ?",
new BeanPropertyRowMapper<>(OtaFirmware.class),
firmwareId
);
return list.isEmpty() ? null : list.get(0);
}
/**
* 按模型查询固件版本列表
*/
public List<OtaFirmware> listFirmwareByModel(Long modelId) {
return jdbcTemplate.query(
"SELECT * FROM iot_ota_firmware WHERE model_id = ? ORDER BY created_at DESC",
new BeanPropertyRowMapper<>(OtaFirmware.class),
modelId
);
}
/**
* 查询所有固件(分页)
*/
public List<OtaFirmware> listAllFirmware(int page, int size) {
int offset = (page - 1) * size;
return jdbcTemplate.query(
"SELECT * FROM iot_ota_firmware ORDER BY created_at DESC LIMIT ? OFFSET ?",
new BeanPropertyRowMapper<>(OtaFirmware.class),
size, offset
);
}
/**
* 设备查询可用固件(基于设备SN查找最新发布的固件)
*/
public OtaFirmware getAvailableFirmware(String deviceSn) {
// 查找设备对应的模型
List<Long> modelIds = jdbcTemplate.queryForList(
"SELECT model_id FROM iot_device WHERE device_sn = ? AND model_id IS NOT NULL",
Long.class, deviceSn
);
if (modelIds.isEmpty()) return null;
Long modelId = modelIds.get(0);
List<OtaFirmware> list = jdbcTemplate.query(
"SELECT * FROM iot_ota_firmware WHERE model_id = ? AND status = 'published' ORDER BY created_at DESC LIMIT 1",
new BeanPropertyRowMapper<>(OtaFirmware.class),
modelId
);
return list.isEmpty() ? null : list.get(0);
}
// ========== 升级任务管理 ==========
/**
* 创建OTA升级任务(按批次)
* @param firmwareId 固件ID
* @param deviceIds 目标设备ID列表
* @param batchSize 每批大小
* @param createdBy 创建人
*/
@Transactional
public OtaTask createUpgradeTask(Long firmwareId, List<Long> deviceIds, int batchSize, String createdBy) {
OtaFirmware firmware = getFirmware(firmwareId);
if (firmware == null) throw new RuntimeException("固件不存在: " + firmwareId);
if (!"published".equals(firmware.getStatus())) throw new RuntimeException("固件未发布: " + firmwareId);
// 序列化设备ID列表
String deviceIdsJson;
try {
deviceIdsJson = objectMapper.writeValueAsString(deviceIds);
} catch (JsonProcessingException e) {
throw new RuntimeException("设备ID序列化失败", e);
}
OtaTask task = new OtaTask();
task.setFirmwareId(firmwareId);
task.setFirmwareVersion(firmware.getFirmwareVersion());
task.setTargetDeviceIds(deviceIdsJson);
task.setTaskStatus("pending");
task.setBatchSize(batchSize > 0 ? batchSize : 10);
task.setTotalDevices(deviceIds.size());
task.setSuccessCount(0);
task.setFailedCount(0);
task.setExecutingCount(0);
task.setCreatedBy(createdBy);
task.setCreatedAt(LocalDateTime.now());
task.setUpdatedAt(LocalDateTime.now());
KeyHolder keyHolder = new GeneratedKeyHolder();
jdbcTemplate.update(connection -> {
var ps = connection.prepareStatement(
"INSERT INTO iot_ota_task (firmware_id, firmware_version, target_device_ids, task_status, " +
"batch_size, total_devices, success_count, failed_count, executing_count, created_by, created_at, updated_at) " +
"VALUES (?, ?, ?::jsonb, ?, ?, ?, 0, 0, 0, ?, NOW(), NOW())",
Statement.RETURN_GENERATED_KEYS
);
ps.setLong(1, firmwareId);
ps.setString(2, firmware.getFirmwareVersion());
ps.setString(3, deviceIdsJson);
ps.setString(4, "pending");
ps.setInt(5, task.getBatchSize());
ps.setInt(6, deviceIds.size());
ps.setString(7, createdBy);
return ps;
}, keyHolder);
task.setId(keyHolder.getKey().longValue());
// 创建升级记录(每台设备一条)
createUpgradeRecords(task.getId(), deviceIds, firmware.getFirmwareVersion());
log.info("Upgrade task created: id={}, firmware={}, devices={}, batchSize={}",
task.getId(), firmware.getFirmwareVersion(), deviceIds.size(), task.getBatchSize());
return task;
}
/**
* 按设备类型/区域创建升级任务
*/
@Transactional
public OtaTask createUpgradeTaskByFilter(Long firmwareId, String deviceType, String area,
int batchSize, String createdBy) {
// 查询符合条件的在线设备
String sql = "SELECT id FROM iot_device WHERE model_id = (SELECT model_id FROM iot_ota_firmware WHERE id = ?) " +
"AND status = 'online'";
List<Object> params = new ArrayList<>();
params.add(firmwareId);
if (deviceType != null && !deviceType.isEmpty()) {
sql += " AND device_type = ?";
params.add(deviceType);
}
if (area != null && !area.isEmpty()) {
sql += " AND area = ?";
params.add(area);
}
List<Long> deviceIds = jdbcTemplate.queryForList(sql, Long.class, params.toArray());
if (deviceIds.isEmpty()) throw new RuntimeException("没有符合条件的在线设备");
OtaTask task = createUpgradeTask(firmwareId, deviceIds, batchSize, createdBy);
task.setTargetType(deviceType);
task.setTargetArea(area);
jdbcTemplate.update(
"UPDATE iot_ota_task SET target_type = ?, target_area = ? WHERE id = ?",
deviceType, area, task.getId()
);
return task;
}
/**
* 启动升级任务(开始执行)
*/
@Transactional
public void startTask(Long taskId) {
OtaTask task = getTask(taskId);
if (task == null) throw new RuntimeException("任务不存在: " + taskId);
if (!"pending".equals(task.getTaskStatus())) throw new RuntimeException("任务状态不允许启动: " + task.getTaskStatus());
jdbcTemplate.update(
"UPDATE iot_ota_task SET task_status = 'executing', executed_at = NOW(), updated_at = NOW() WHERE id = ?",
taskId
);
// 将第一批次的记录设为executing
int batchSize = task.getBatchSize();
jdbcTemplate.update(
"UPDATE iot_ota_upgrade_record SET status = 'executing', started_at = NOW() " +
"WHERE task_id = ? AND id IN (SELECT id FROM iot_ota_upgrade_record WHERE task_id = ? AND status = 'pending' LIMIT ?)",
taskId, taskId, batchSize
);
log.info("Task {} started, first batch size: {}", taskId, batchSize);
}
/**
* 取消升级任务
*/
@Transactional
public void cancelTask(Long taskId) {
jdbcTemplate.update(
"UPDATE iot_ota_task SET task_status = 'cancelled', updated_at = NOW() WHERE id = ? AND task_status IN ('pending', 'executing')",
taskId
);
// 取消未完成的记录
jdbcTemplate.update(
"UPDATE iot_ota_upgrade_record SET status = 'failed', fail_reason = '任务已取消' " +
"WHERE task_id = ? AND status IN ('pending', 'executing')",
taskId
);
}
/**
* 查询升级任务详情
*/
public OtaTask getTask(Long taskId) {
List<OtaTask> list = jdbcTemplate.query(
"SELECT * FROM iot_ota_task WHERE id = ?",
new BeanPropertyRowMapper<>(OtaTask.class),
taskId
);
return list.isEmpty() ? null : list.get(0);
}
/**
* 查询升级任务列表
*/
public List<OtaTask> listTasks(int page, int size) {
int offset = (page - 1) * size;
return jdbcTemplate.query(
"SELECT * FROM iot_ota_task ORDER BY created_at DESC LIMIT ? OFFSET ?",
new BeanPropertyRowMapper<>(OtaTask.class),
size, offset
);
}
// ========== 进度追踪 ==========
/**
* 更新单台设备升级进度
*/
@Transactional
public void updateProgress(Long recordId, int progress) {
jdbcTemplate.update(
"UPDATE iot_ota_upgrade_record SET progress = ? WHERE id = ?",
progress, recordId
);
}
/**
* 标记设备升级成功
*/
@Transactional
public void markSuccess(Long recordId, Long deviceId, String toVersion) {
jdbcTemplate.update(
"UPDATE iot_ota_upgrade_record SET status = 'success', progress = 100, completed_at = NOW() WHERE id = ?",
recordId
);
jdbcTemplate.update(
"UPDATE iot_ota_task SET success_count = success_count + 1, executing_count = executing_count - 1, updated_at = NOW() WHERE id = (SELECT task_id FROM iot_ota_upgrade_record WHERE id = ?)",
recordId
);
// 更新设备固件版本
jdbcTemplate.update(
"UPDATE iot_device SET firmware_version = ? WHERE id = ?",
toVersion, deviceId
);
// 检查是否触发下一批次
triggerNextBatch(recordId);
}
/**
* 标记设备升级失败
*/
@Transactional
public void markFailed(Long recordId, String failReason) {
jdbcTemplate.update(
"UPDATE iot_ota_upgrade_record SET status = 'failed', fail_reason = ?, completed_at = NOW() WHERE id = ?",
failReason, recordId
);
jdbcTemplate.update(
"UPDATE iot_ota_task SET failed_count = failed_count + 1, executing_count = executing_count - 1, updated_at = NOW() WHERE id = (SELECT task_id FROM iot_ota_upgrade_record WHERE id = ?)",
recordId
);
triggerNextBatch(recordId);
}
/**
* 查询升级任务统计
*/
public Map<String, Object> getTaskStatistics(Long taskId) {
OtaTask task = getTask(taskId);
if (task == null) return Collections.emptyMap();
Map<String, Object> stats = new LinkedHashMap<>();
stats.put("taskId", taskId);
stats.put("taskStatus", task.getTaskStatus());
stats.put("totalDevices", task.getTotalDevices());
stats.put("successCount", task.getSuccessCount());
stats.put("failedCount", task.getFailedCount());
stats.put("executingCount", task.getExecutingCount());
stats.put("pendingCount", task.getTotalDevices() - task.getSuccessCount() - task.getFailedCount() - task.getExecutingCount());
double progress = task.getTotalDevices() > 0 ?
((task.getSuccessCount() + task.getFailedCount()) * 100.0 / task.getTotalDevices()) : 0;
stats.put("progressPercent", Math.round(progress * 100) / 100.0);
// 按状态分组统计
List<Map<String, Object>> byStatus = jdbcTemplate.queryForList(
"SELECT status, COUNT(*) as count FROM iot_ota_upgrade_record WHERE task_id = ? GROUP BY status",
taskId
);
stats.put("statusBreakdown", byStatus);
return stats;
}
/**
* 查询升级记录列表
*/
public List<OtaUpgradeRecord> getTaskRecords(Long taskId) {
return jdbcTemplate.query(
"SELECT * FROM iot_ota_upgrade_record WHERE task_id = ? ORDER BY id",
new BeanPropertyRowMapper<>(OtaUpgradeRecord.class),
taskId
);
}
/**
* 查询单台设备的升级历史
*/
public List<OtaUpgradeRecord> getDeviceUpgradeHistory(String deviceSn) {
return jdbcTemplate.query(
"SELECT r.* FROM iot_ota_upgrade_record r JOIN iot_device d ON r.device_id = d.id " +
"WHERE d.device_sn = ? ORDER BY r.created_at DESC",
new BeanPropertyRowMapper<>(OtaUpgradeRecord.class),
deviceSn
);
}
// ===== 私有辅助方法 =====
/**
* 创建升级记录
*/
private void createUpgradeRecords(Long taskId, List<Long> deviceIds, String toVersion) {
String sql = "INSERT INTO iot_ota_upgrade_record (task_id, device_id, device_sn, from_version, to_version, status, progress, created_at) " +
"VALUES (?, ?, ?, ?, ?, 'pending', 0, NOW())";
List<Object[]> batchArgs = new ArrayList<>();
for (Long deviceId : deviceIds) {
// 查询设备SN和当前版本
List<Map<String, Object>> deviceInfo = jdbcTemplate.queryForList(
"SELECT device_sn, firmware_version FROM iot_device WHERE id = ?",
deviceId
);
if (!deviceInfo.isEmpty()) {
Map<String, Object> info = deviceInfo.get(0);
String deviceSn = (String) info.get("device_sn");
String fromVersion = info.get("firmware_version") != null ?
(String) info.get("firmware_version") : "unknown";
batchArgs.add(new Object[]{taskId, deviceId, deviceSn, fromVersion, toVersion});
}
}
if (!batchArgs.isEmpty()) {
jdbcTemplate.batchUpdate(sql, batchArgs);
}
}
/**
* 触发下一批次
* 当一个批次的设备完成(成功或失败)后,自动启动下一批
*/
private void triggerNextBatch(Long recordId) {
Long taskId = jdbcTemplate.queryForObject(
"SELECT task_id FROM iot_ota_upgrade_record WHERE id = ?", Long.class, recordId
);
OtaTask task = getTask(taskId);
if (task == null || !"executing".equals(task.getTaskStatus())) return;
// 检查是否还有执行中的记录
int executing = jdbcTemplate.queryForObject(
"SELECT COUNT(*) FROM iot_ota_upgrade_record WHERE task_id = ? AND status = 'executing'",
Integer.class, taskId
);
if (executing == 0) {
// 启动下一批
int pending = jdbcTemplate.queryForObject(
"SELECT COUNT(*) FROM iot_ota_upgrade_record WHERE task_id = ? AND status = 'pending'",
Integer.class, taskId
);
if (pending > 0) {
int nextBatch = Math.min(pending, task.getBatchSize());
jdbcTemplate.update(
"UPDATE iot_ota_upgrade_record SET status = 'executing', started_at = NOW() " +
"WHERE task_id = ? AND id IN (SELECT id FROM iot_ota_upgrade_record WHERE task_id = ? AND status = 'pending' LIMIT ?)",
taskId, taskId, nextBatch
);
jdbcTemplate.update(
"UPDATE iot_ota_task SET executing_count = executing_count + ? WHERE id = ?",
nextBatch, taskId
);
log.info("Task {} next batch started: {} devices", taskId, nextBatch);
} else {
// 所有设备处理完毕,标记任务完成
int failedCount = jdbcTemplate.queryForObject(
"SELECT COUNT(*) FROM iot_ota_upgrade_record WHERE task_id = ? AND status = 'failed'",
Integer.class, taskId
);
String finalStatus = failedCount > 0 ? "completed" : "completed";
jdbcTemplate.update(
"UPDATE iot_ota_task SET task_status = ?, completed_at = NOW(), updated_at = NOW() WHERE id = ?",
finalStatus, taskId
);
log.info("Task {} completed: total={}, success={}, failed={}",
taskId, task.getTotalDevices(), task.getSuccessCount(), failedCount);
}
}
}
}
@@ -0,0 +1,84 @@
-- =============================================
-- 智慧水务管理系统 - 设备影子 + OTA固件升级 DDL
-- 版本: V_shadow_ota
-- =============================================
-- OTA固件版本表
CREATE TABLE IF NOT EXISTS iot_ota_firmware (
id BIGSERIAL PRIMARY KEY,
model_id BIGINT NOT NULL REFERENCES iot_device_model(id),
firmware_version VARCHAR(20) NOT NULL,
file_url VARCHAR(500) NOT NULL,
description TEXT,
status VARCHAR(20) NOT NULL DEFAULT 'draft', -- draft/published/deprecated
md5 VARCHAR(64),
file_size BIGINT DEFAULT 0,
published_by VARCHAR(50),
published_at TIMESTAMP,
created_at TIMESTAMP DEFAULT NOW(),
updated_at TIMESTAMP DEFAULT NOW(),
UNIQUE(model_id, firmware_version)
);
COMMENT ON TABLE iot_ota_firmware IS 'OTA固件版本表';
COMMENT ON COLUMN iot_ota_firmware.model_id IS '关联设备模型ID';
COMMENT ON COLUMN iot_ota_firmware.firmware_version IS '固件版本号';
COMMENT ON COLUMN iot_ota_firmware.file_url IS '固件文件下载地址';
COMMENT ON COLUMN iot_ota_firmware.status IS '状态: draft-草稿/published-已发布/deprecated-已废弃';
COMMENT ON COLUMN iot_ota_firmware.md5 IS 'MD5校验值';
CREATE INDEX IF NOT EXISTS idx_ota_firmware_model ON iot_ota_firmware(model_id);
CREATE INDEX IF NOT EXISTS idx_ota_firmware_status ON iot_ota_firmware(status);
-- OTA升级任务表
CREATE TABLE IF NOT EXISTS iot_ota_task (
id BIGSERIAL PRIMARY KEY,
firmware_id BIGINT NOT NULL REFERENCES iot_ota_firmware(id),
firmware_version VARCHAR(20) NOT NULL,
target_type VARCHAR(30), -- 目标设备类型
target_area VARCHAR(50), -- 目标区域
target_device_ids JSONB, -- 目标设备ID列表
task_status VARCHAR(20) NOT NULL DEFAULT 'pending', -- pending/executing/completed/failed/cancelled
batch_size INT DEFAULT 10, -- 每批升级数量
total_devices INT DEFAULT 0, -- 总设备数
success_count INT DEFAULT 0,
failed_count INT DEFAULT 0,
executing_count INT DEFAULT 0,
created_by VARCHAR(50),
executed_at TIMESTAMP,
completed_at TIMESTAMP,
created_at TIMESTAMP DEFAULT NOW(),
updated_at TIMESTAMP DEFAULT NOW()
);
COMMENT ON TABLE iot_ota_task IS 'OTA升级任务表';
COMMENT ON COLUMN iot_ota_task.task_status IS '任务状态: pending-待执行/executing-执行中/completed-已完成/failed-失败/cancelled-已取消';
COMMENT ON COLUMN iot_ota_task.batch_size IS '每批升级设备数量';
CREATE INDEX IF NOT EXISTS idx_ota_task_firmware ON iot_ota_task(firmware_id);
CREATE INDEX IF NOT EXISTS idx_ota_task_status ON iot_ota_task(task_status);
CREATE INDEX IF NOT EXISTS idx_ota_task_created ON iot_ota_task(created_at DESC);
-- OTA升级记录表(每台设备一条记录)
CREATE TABLE IF NOT EXISTS iot_ota_upgrade_record (
id BIGSERIAL PRIMARY KEY,
task_id BIGINT NOT NULL REFERENCES iot_ota_task(id),
device_id BIGINT NOT NULL REFERENCES iot_device(id),
device_sn VARCHAR(100) NOT NULL,
from_version VARCHAR(20), -- 升级前版本
to_version VARCHAR(20), -- 目标版本
status VARCHAR(20) NOT NULL DEFAULT 'pending', -- pending/executing/success/failed/timeout
fail_reason VARCHAR(500),
progress INT DEFAULT 0, -- 进度百分比 0-100
started_at TIMESTAMP,
completed_at TIMESTAMP,
created_at TIMESTAMP DEFAULT NOW()
);
COMMENT ON TABLE iot_ota_upgrade_record IS 'OTA升级记录表(单台设备)';
COMMENT ON COLUMN iot_ota_upgrade_record.status IS '升级状态: pending-待升级/executing-升级中/success-成功/failed-失败/timeout-超时';
CREATE INDEX IF NOT EXISTS idx_ota_record_task ON iot_ota_upgrade_record(task_id);
CREATE INDEX IF NOT EXISTS idx_ota_record_device ON iot_ota_upgrade_record(device_id);
CREATE INDEX IF NOT EXISTS idx_ota_record_sn ON iot_ota_upgrade_record(device_sn);
CREATE INDEX IF NOT EXISTS idx_ota_record_status ON iot_ota_upgrade_record(status);
@@ -0,0 +1,195 @@
package com.water.iot.service;
import com.fasterxml.jackson.databind.ObjectMapper;
import com.water.iot.entity.DeviceShadow;
import org.junit.jupiter.api.BeforeEach;
import org.junit.jupiter.api.Test;
import org.junit.jupiter.api.extension.ExtendWith;
import org.mockito.Mock;
import org.mockito.junit.jupiter.MockitoExtension;
import org.springframework.data.redis.core.HashOperations;
import org.springframework.data.redis.core.StringRedisTemplate;
import org.springframework.jdbc.core.JdbcTemplate;
import java.util.*;
import static org.junit.jupiter.api.Assertions.*;
import static org.mockito.ArgumentMatchers.*;
import static org.mockito.Mockito.*;
/**
* DeviceShadowService 单元测试
* Mock Redis + Mock JdbcTemplate
*/
@ExtendWith(MockitoExtension.class)
class DeviceShadowServiceTest {
@Mock
private StringRedisTemplate redisTemplate;
@Mock
private JdbcTemplate jdbcTemplate;
@Mock
private HashOperations<String, Object, Object> hashOperations;
private DeviceShadowService shadowService;
@BeforeEach
void setUp() {
when(redisTemplate.opsForHash()).thenReturn(hashOperations);
shadowService = new DeviceShadowService(redisTemplate, jdbcTemplate, new ObjectMapper());
}
@Test
void updateReported_shouldUpdateRedisAndDb() {
// Given
String deviceSn = "FM-001";
Map<String, Object> state = Map.of("flow", 12.5, "pressure", 0.35);
when(hashOperations.get(eq("shadow:" + deviceSn), eq("version"))).thenReturn("1");
when(hashOperations.get(eq("shadow:" + deviceSn), eq("desired"))).thenReturn(null);
when(jdbcTemplate.update(anyString(), any())).thenReturn(1);
when(jdbcTemplate.queryForList(anyString(), eq(Long.class), anyString()))
.thenReturn(List.of(1L));
// When
shadowService.updateReported(deviceSn, state);
// Then
verify(hashOperations).putAll(eq("shadow:" + deviceSn), anyMap());
verify(redisTemplate).expire(eq("shadow:" + deviceSn), eq(24L), any());
verify(jdbcTemplate).update(contains("UPDATE iot_device"), eq(deviceSn));
}
@Test
void updateDesired_shouldComputeDelta() {
// Given
String deviceSn = "FM-001";
Map<String, Object> desired = Map.of("valveOpen", 80, "alarm", true);
when(hashOperations.get(eq("shadow:" + deviceSn), eq("version"))).thenReturn("2");
when(hashOperations.get(eq("shadow:" + deviceSn), eq("reported")))
.thenReturn("{\"valveOpen\":50,\"alarm\":true}");
when(jdbcTemplate.update(anyString(), any())).thenReturn(1);
when(jdbcTemplate.queryForList(anyString(), eq(Long.class), anyString()))
.thenReturn(List.of(1L));
// When
shadowService.updateDesired(deviceSn, desired);
// Then
verify(hashOperations).putAll(eq("shadow:" + deviceSn), argThat(map -> {
String delta = (String) map.get("delta");
// delta should contain valveOpen:80 (different from reported 50)
// but NOT alarm (same as reported)
return delta != null && delta.contains("valveOpen") && delta.contains("80");
}));
}
@Test
void getShadow_fromRedis_shouldReturnShadow() {
// Given
String deviceSn = "FM-001";
Map<Object, Object> entries = new HashMap<>();
entries.put("reported", "{\"flow\":12.5}");
entries.put("desired", "{\"flow\":15.0}");
entries.put("delta", "{\"flow\":15.0}");
entries.put("version", "3");
entries.put("online", "true");
entries.put("lastReport", "2026-06-14T10:00:00");
when(hashOperations.entries("shadow:" + deviceSn)).thenReturn(entries);
// When
DeviceShadow shadow = shadowService.getShadow(deviceSn);
// Then
assertNotNull(shadow);
assertEquals(deviceSn, shadow.getDeviceSn());
assertEquals("{\"flow\":12.5}", shadow.getReportedState());
assertEquals("{\"flow\":15.0}", shadow.getDesiredState());
assertEquals(3L, shadow.getVersion());
assertEquals("online", shadow.getOnlineStatus());
}
@Test
void getDesired_emptyRedis_shouldReturnEmpty() {
when(hashOperations.get(anyString(), eq("desired"))).thenReturn(null);
Map<String, Object> result = shadowService.getDesired("FM-001");
assertTrue(result.isEmpty());
}
@Test
void isOnline_true() {
when(hashOperations.get("shadow:FM-001", "online")).thenReturn("true");
assertTrue(shadowService.isOnline("FM-001"));
}
@Test
void isOnline_false() {
when(hashOperations.get("shadow:FM-001", "online")).thenReturn("false");
assertFalse(shadowService.isOnline("FM-001"));
}
@Test
void batchGetShadows_shouldReturnMultiple() {
// Given
List<String> sns = List.of("FM-001", "FM-002");
Map<Object, Object> entries1 = new HashMap<>();
entries1.put("reported", "{\"flow\":12}");
entries1.put("version", "1");
entries1.put("online", "true");
Map<Object, Object> entries2 = new HashMap<>();
entries2.put("reported", "{\"pressure\":0.4}");
entries2.put("version", "2");
entries2.put("online", "false");
when(hashOperations.entries("shadow:FM-001")).thenReturn(entries1);
when(hashOperations.entries("shadow:FM-002")).thenReturn(entries2);
// When
List<DeviceShadow> shadows = shadowService.batchGetShadows(sns);
// Then
assertEquals(2, shadows.size());
}
@Test
void detectOfflineDevices_shouldMarkOffline() {
// Given
Set<String> keys = Set.of("shadow:FM-001", "shadow:FM-002");
when(redisTemplate.keys("shadow:*")).thenReturn(keys);
// FM-001: last report 60 minutes ago (offline)
when(hashOperations.get(eq("shadow:FM-001"), eq("lastReport")))
.thenReturn(java.time.LocalDateTime.now().minusMinutes(60).toString());
// FM-002: last report 5 minutes ago (online)
when(hashOperations.get(eq("shadow:FM-002"), eq("lastReport")))
.thenReturn(java.time.LocalDateTime.now().minusMinutes(5).toString());
when(jdbcTemplate.update(anyString(), anyString())).thenReturn(1);
// When
List<String> offline = shadowService.detectOfflineDevices(30);
// Then
assertEquals(1, offline.size());
assertEquals("FM-001", offline.get(0));
verify(hashOperations).put("shadow:FM-001", "online", "false");
}
@Test
void deleteShadow_shouldRemoveFromRedisAndDb() {
when(redisTemplate.delete("shadow:FM-001")).thenReturn(true);
when(jdbcTemplate.queryForList(anyString(), eq(Long.class), anyString()))
.thenReturn(List.of(1L));
when(jdbcTemplate.update(anyString(), anyLong())).thenReturn(1);
shadowService.deleteShadow("FM-001");
verify(redisTemplate).delete("shadow:FM-001");
verify(jdbcTemplate).update(contains("DELETE FROM iot_device_shadow"), eq(1L));
}
}
@@ -0,0 +1,310 @@
package com.water.iot.service;
import com.fasterxml.jackson.databind.ObjectMapper;
import com.water.iot.entity.OtaFirmware;
import com.water.iot.entity.OtaTask;
import com.water.iot.entity.OtaUpgradeRecord;
import org.junit.jupiter.api.BeforeEach;
import org.junit.jupiter.api.Test;
import org.junit.jupiter.api.extension.ExtendWith;
import org.mockito.Mock;
import org.mockito.junit.jupiter.MockitoExtension;
import org.springframework.jdbc.core.BeanPropertyRowMapper;
import org.springframework.jdbc.core.JdbcTemplate;
import org.springframework.jdbc.support.GeneratedKeyHolder;
import java.util.*;
import static org.junit.jupiter.api.Assertions.*;
import static org.mockito.ArgumentMatchers.*;
import static org.mockito.Mockito.*;
/**
* OtaService 单元测试
* Mock JdbcTemplate
*/
@ExtendWith(MockitoExtension.class)
class OtaServiceTest {
@Mock
private JdbcTemplate jdbcTemplate;
@Mock
private DeviceShadowService shadowService;
private OtaService otaService;
private final ObjectMapper objectMapper = new ObjectMapper();
@BeforeEach
void setUp() {
otaService = new OtaService(jdbcTemplate, shadowService, objectMapper);
}
@Test
void createFirmware_shouldSetDraftStatusAndReturn() {
OtaFirmware firmware = new OtaFirmware();
firmware.setModelId(1L);
firmware.setFirmwareVersion("v2.0.0");
firmware.setFileUrl("http://oss.example.com/firmware/v2.0.0.bin");
firmware.setDescription("压力传感器固件升级");
firmware.setMd5("d41d8cd98f00b204e9800998ecf8427e");
firmware.setFileSize(1024000L);
when(jdbcTemplate.update(any(org.springframework.jdbc.core.PreparedStatementCreator.class),
any(GeneratedKeyHolder.class))).thenReturn(1);
OtaFirmware result = otaService.createFirmware(firmware);
assertNotNull(result);
assertEquals("draft", result.getStatus());
assertEquals("v2.0.0", result.getFirmwareVersion());
verify(jdbcTemplate).update(any(org.springframework.jdbc.core.PreparedStatementCreator.class),
any(GeneratedKeyHolder.class));
}
@Test
void publishFirmware_shouldUpdateStatus() {
when(jdbcTemplate.update(anyString(), anyString(), anyLong())).thenReturn(1);
otaService.publishFirmware(1L, "admin");
verify(jdbcTemplate).update(contains("UPDATE iot_ota_firmware"), eq("admin"), eq(1L));
}
@Test
void publishFirmware_notFound_shouldThrow() {
when(jdbcTemplate.update(anyString(), anyString(), anyLong())).thenReturn(0);
assertThrows(RuntimeException.class, () -> otaService.publishFirmware(999L, "admin"));
}
@Test
void getFirmware_shouldReturnFirmware() {
OtaFirmware fw = new OtaFirmware();
fw.setId(1L);
fw.setFirmwareVersion("v2.0.0");
when(jdbcTemplate.query(anyString(), any(BeanPropertyRowMapper.class), anyLong()))
.thenReturn(List.of(fw));
OtaFirmware result = otaService.getFirmware(1L);
assertNotNull(result);
assertEquals("v2.0.0", result.getFirmwareVersion());
}
@Test
void getFirmware_notFound_shouldReturnNull() {
when(jdbcTemplate.query(anyString(), any(BeanPropertyRowMapper.class), anyLong()))
.thenReturn(Collections.emptyList());
assertNull(otaService.getFirmware(999L));
}
@Test
void listAllFirmware_shouldReturnPaged() {
OtaFirmware fw1 = new OtaFirmware();
fw1.setId(1L);
OtaFirmware fw2 = new OtaFirmware();
fw2.setId(2L);
when(jdbcTemplate.query(anyString(), any(BeanPropertyRowMapper.class), anyInt(), anyInt()))
.thenReturn(List.of(fw1, fw2));
List<OtaFirmware> result = otaService.listAllFirmware(1, 10);
assertEquals(2, result.size());
}
@Test
void getAvailableFirmware_shouldReturnLatestPublished() {
// Device has model_id = 1
when(jdbcTemplate.queryForList(anyString(), eq(Long.class), anyString()))
.thenReturn(List.of(1L));
OtaFirmware fw = new OtaFirmware();
fw.setFirmwareVersion("v2.0.0");
fw.setStatus("published");
when(jdbcTemplate.query(contains("published"), any(BeanPropertyRowMapper.class), anyLong()))
.thenReturn(List.of(fw));
OtaFirmware result = otaService.getAvailableFirmware("FM-001");
assertNotNull(result);
assertEquals("v2.0.0", result.getFirmwareVersion());
}
@Test
void getAvailableFirmware_noModel_shouldReturnNull() {
when(jdbcTemplate.queryForList(anyString(), eq(Long.class), anyString()))
.thenReturn(Collections.emptyList());
assertNull(otaService.getAvailableFirmware("FM-001"));
}
@Test
void createUpgradeTask_shouldCreateTaskAndRecords() {
// Mock firmware exists and is published
OtaFirmware fw = new OtaFirmware();
fw.setId(1L);
fw.setFirmwareVersion("v2.0.0");
fw.setStatus("published");
when(jdbcTemplate.query(anyString(), any(BeanPropertyRowMapper.class), anyLong()))
.thenReturn(List.of(fw));
// Mock insert task
when(jdbcTemplate.update(any(org.springframework.jdbc.core.PreparedStatementCreator.class),
any(GeneratedKeyHolder.class))).thenReturn(1);
// Mock device info query
Map<String, Object> deviceInfo = new HashMap<>();
deviceInfo.put("device_sn", "FM-001");
deviceInfo.put("firmware_version", "v1.0.0");
when(jdbcTemplate.queryForList(contains("SELECT device_sn"), anyLong()))
.thenReturn(List.of(deviceInfo));
// Mock batch insert
when(jdbcTemplate.batchUpdate(anyString(), anyList())).thenReturn(new int[]{1});
OtaTask task = otaService.createUpgradeTask(1L, List.of(1L, 2L), 5, "admin");
assertNotNull(task);
assertEquals("v2.0.0", task.getFirmwareVersion());
assertEquals("pending", task.getTaskStatus());
assertEquals(2, task.getTotalDevices());
assertEquals(5, task.getBatchSize());
}
@Test
void createUpgradeTask_unpublishedFirmware_shouldThrow() {
OtaFirmware fw = new OtaFirmware();
fw.setId(1L);
fw.setStatus("draft");
when(jdbcTemplate.query(anyString(), any(BeanPropertyRowMapper.class), anyLong()))
.thenReturn(List.of(fw));
assertThrows(RuntimeException.class,
() -> otaService.createUpgradeTask(1L, List.of(1L), 5, "admin"));
}
@Test
void startTask_shouldUpdateStatusAndFirstBatch() {
// Mock task exists
OtaTask task = new OtaTask();
task.setId(1L);
task.setTaskStatus("pending");
task.setBatchSize(5);
when(jdbcTemplate.query(anyString(), any(BeanPropertyRowMapper.class), anyLong()))
.thenReturn(List.of(task));
when(jdbcTemplate.update(anyString(), any())).thenReturn(1);
otaService.startTask(1L);
verify(jdbcTemplate, atLeast(2)).update(anyString(), any());
}
@Test
void cancelTask_shouldUpdateStatusAndRecords() {
when(jdbcTemplate.update(anyString(), anyLong())).thenReturn(1);
otaService.cancelTask(1L);
verify(jdbcTemplate, times(2)).update(anyString(), anyLong());
}
@Test
void getTaskStatistics_shouldReturnStats() {
OtaTask task = new OtaTask();
task.setId(1L);
task.setTaskStatus("executing");
task.setTotalDevices(10);
task.setSuccessCount(5);
task.setFailedCount(1);
task.setExecutingCount(2);
when(jdbcTemplate.query(anyString(), any(BeanPropertyRowMapper.class), anyLong()))
.thenReturn(List.of(task));
List<Map<String, Object>> statusBreakdown = List.of(
Map.of("status", "success", "count", 5L),
Map.of("status", "failed", "count", 1L),
Map.of("status", "executing", "count", 2L)
);
when(jdbcTemplate.queryForList(anyString(), anyLong())).thenReturn(statusBreakdown);
Map<String, Object> stats = otaService.getTaskStatistics(1L);
assertNotNull(stats);
assertEquals(10, stats.get("totalDevices"));
assertEquals(5, stats.get("successCount"));
assertEquals(60.0, stats.get("progressPercent"));
assertNotNull(stats.get("statusBreakdown"));
}
@Test
void getTaskRecords_shouldReturnRecords() {
OtaUpgradeRecord record = new OtaUpgradeRecord();
record.setId(1L);
record.setTaskId(1L);
record.setStatus("success");
when(jdbcTemplate.query(anyString(), any(BeanPropertyRowMapper.class), anyLong()))
.thenReturn(List.of(record));
List<OtaUpgradeRecord> records = otaService.getTaskRecords(1L);
assertEquals(1, records.size());
assertEquals("success", records.get(0).getStatus());
}
@Test
void markSuccess_shouldUpdateRecordAndTask() {
when(jdbcTemplate.update(anyString(), any())).thenReturn(1);
when(jdbcTemplate.queryForObject(anyString(), eq(Long.class), anyLong())).thenReturn(1L);
// Mock task for triggerNextBatch
OtaTask task = new OtaTask();
task.setId(1L);
task.setTaskStatus("executing");
task.setBatchSize(5);
when(jdbcTemplate.query(anyString(), any(BeanPropertyRowMapper.class), anyLong()))
.thenReturn(List.of(task));
when(jdbcTemplate.queryForObject(contains("COUNT"), eq(Integer.class), anyLong()))
.thenReturn(0, 3); // 0 executing, 3 pending
otaService.markSuccess(1L, 1L, "v2.0.0");
verify(jdbcTemplate, atLeast(3)).update(anyString(), any());
}
@Test
void markFailed_shouldUpdateWithReason() {
when(jdbcTemplate.update(anyString(), any())).thenReturn(1);
when(jdbcTemplate.queryForObject(anyString(), eq(Long.class), anyLong())).thenReturn(1L);
OtaTask task = new OtaTask();
task.setId(1L);
task.setTaskStatus("executing");
task.setBatchSize(5);
when(jdbcTemplate.query(anyString(), any(BeanPropertyRowMapper.class), anyLong()))
.thenReturn(List.of(task));
when(jdbcTemplate.queryForObject(contains("COUNT"), eq(Integer.class), anyLong()))
.thenReturn(0, 0); // 0 executing, 0 pending -> task completed
otaService.markFailed(1L, "设备连接超时");
verify(jdbcTemplate).update(contains("fail_reason"), eq("设备连接超时"), eq(1L));
}
@Test
void updateProgress_shouldUpdateRecord() {
when(jdbcTemplate.update(anyString(), anyInt(), anyLong())).thenReturn(1);
otaService.updateProgress(1L, 50);
verify(jdbcTemplate).update(contains("progress"), eq(50), eq(1L));
}
@Test
void deprecateFirmware_shouldUpdateStatus() {
when(jdbcTemplate.update(anyString(), anyLong())).thenReturn(1);
otaService.deprecateFirmware(1L);
verify(jdbcTemplate).update(contains("deprecated"), eq(1L));
}
@Test
void listTasks_shouldReturnPaged() {
OtaTask t1 = new OtaTask();
t1.setId(1L);
OtaTask t2 = new OtaTask();
t2.setId(2L);
when(jdbcTemplate.query(anyString(), any(BeanPropertyRowMapper.class), anyInt(), anyInt()))
.thenReturn(List.of(t1, t2));
List<OtaTask> result = otaService.listTasks(1, 10);
assertEquals(2, result.size());
}
}