@@ -11,6 +11,7 @@ import lombok.Data;
import lombok.extern.slf4j.Slf4j ;
import org.springframework.beans.BeanUtils ;
import org.springframework.beans.factory.annotation.Autowired ;
import org.springframework.http.HttpStatus ;
import org.springframework.http.ResponseEntity ;
import org.springframework.web.bind.annotation.* ;
import org.springframework.web.multipart.MultipartFile ;
@@ -20,6 +21,16 @@ import jakarta.validation.constraints.NotBlank;
import java.time.LocalDateTime ;
import java.util.HashMap ;
import java.util.Map ;
import org.springframework.http.HttpEntity ;
import org.springframework.http.HttpHeaders ;
import org.springframework.http.HttpMethod ;
import org.springframework.http.MediaType ;
import org.springframework.web.client.RestTemplate ;
import com.fasterxml.jackson.databind.ObjectMapper ;
import com.fasterxml.jackson.databind.JsonNode ;
import com.rj.service.MinIOService ;
import java.io.* ;
import java.net.URL ;
/**
* 图像模型控制器
@@ -49,6 +60,17 @@ public class ImageModelController {
*/
@Autowired
private IImageModelService imageModelService ;
@Autowired
private RestTemplate restTemplate ;
@Autowired
private com . rj . scheduler . ImageModelStatusScheduler imageModelStatusScheduler ;
@Autowired
private MinIOService minIOService ;
private final ObjectMapper objectMapper = new ObjectMapper ( ) ;
/**
* 图像模型请求DTO
@@ -649,4 +671,496 @@ public class ImageModelController {
return ResponseEntity . status ( 500 ) . body ( response ) ;
}
}
/**
* 图像处理接口
*/
@PostMapping ( " /image-background " )
@Operation ( summary = " 图像处理 " , description = " 基于参考图像和提示词进行图像处理 " )
public ResponseEntity < Map < String , Object > > imageProcessing (
@Parameter ( description = " 图像处理请求参数 " ) @Valid @RequestBody TextModelController . ImageProcessingRequest request ) {
log . info ( " 开始图像处理,request : {} " , request ) ;
Map < String , Object > result = new HashMap < > ( ) ;
try {
log . info ( " 开始图像处理,基础图片: {}, 参考图片: {}, 参考提示词: {} " ,
request . getBaseImageUrl ( ) , request . getRefImageUrl ( ) , request . getRefPrompt ( ) ) ;
// 模拟图像处理(实际项目中应该调用图像处理服务)
Map < String , Object > processingResponse = performImageProcessing ( request ) ;
result . put ( " success " , true ) ;
result . put ( " message " , " 图像处理完成 " ) ;
result . put ( " baseImageUrl " , request . getBaseImageUrl ( ) ) ;
result . put ( " refImageUrl " , request . getRefImageUrl ( ) ) ;
result . put ( " refPrompt " , request . getRefPrompt ( ) ) ;
result . put ( " processedImages " , processingResponse . get ( " processedImages " ) ) ;
result . put ( " processingTimeMs " , processingResponse . get ( " processingTimeMs " ) ) ;
result . put ( " modelVersion " , request . getParameters ( ) ! = null ? request . getParameters ( ) . getModelVersion ( ) : " v3 " ) ;
result . put ( " refPromptWeight " , request . getParameters ( ) ! = null ? request . getParameters ( ) . getRefPromptWeight ( ) : 0 . 5 ) ;
result . put ( " requestId " , generateRequestId ( ) ) ;
result . put ( " timestamp " , LocalDateTime . now ( ) ) ;
log . info ( " 图像处理完成,生成图片数量: {} " , processingResponse . get ( " imageCount " ) ) ;
return ResponseEntity . ok ( result ) ;
} catch ( Exception e ) {
log . error ( " 图像处理失败: {} " , e . getMessage ( ) , e ) ;
result . put ( " success " , false ) ;
result . put ( " message " , " 图像处理失败: " + e . getMessage ( ) ) ;
result . put ( " error " , e . getClass ( ) . getSimpleName ( ) ) ;
return ResponseEntity . status ( HttpStatus . INTERNAL_SERVER_ERROR ) . body ( result ) ;
}
}
/**
* 执行图像处理的方法 - 调用阿里云API
*/
private Map < String , Object > performImageProcessing ( TextModelController . ImageProcessingRequest request ) {
Map < String , Object > response = new HashMap < > ( ) ;
long startTime = System . currentTimeMillis ( ) ;
try {
log . info ( " 开始调用阿里云图像生成API, 基础图片: {}, 参考图片: {} " ,
request . getBaseImageUrl ( ) , request . getRefImageUrl ( ) ) ;
// 处理引导图像RGBA转换( 如果存在引导图像)
if ( request . getRefImageUrl ( ) ! = null & & ! request . getRefImageUrl ( ) . trim ( ) . isEmpty ( ) ) {
log . info ( " 检测到引导图像, 开始进行RGBA转换处理 " ) ;
String processedRefImageUrl = processRefImageForRGBA ( request . getRefImageUrl ( ) ) ;
if ( processedRefImageUrl ! = null ) {
log . info ( " 引导图像RGBA转换完成, 新URL: {} " , processedRefImageUrl ) ;
request . setRefImageUrl ( processedRefImageUrl ) ;
}
}
// 构建阿里云请求参数
Map < String , Object > aliyunRequest = buildAliyunRequest ( request ) ;
// 发送请求到阿里云
String aliyunResponse = sendRequestToAliyun ( aliyunRequest ) ;
// 解析阿里云响应
Map < String , Object > parsedResponse = parseAliyunResponse ( aliyunResponse ) ;
// 保存数据到数据库
ImageModel imageModel = createImageModelRecord ( request , parsedResponse ) ;
boolean saveResult = imageModelService . saveImageModel ( imageModel ) ;
long processingTime = System . currentTimeMillis ( ) - startTime ;
response . put ( " success " , true ) ;
response . put ( " taskId " , parsedResponse . get ( " taskId " ) ) ;
response . put ( " status " , parsedResponse . get ( " status " ) ) ;
response . put ( " taskStatus " , parsedResponse . get ( " status " ) ) ; // 添加task_status字段
response . put ( " processingTimeMs " , processingTime ) ;
response . put ( " saveResult " , saveResult ) ;
response . put ( " recordId " , imageModel . getUuid ( ) ) ;
// 如果有参考边缘信息,记录处理详情
if ( request . getReferenceEdge ( ) ! = null ) {
response . put ( " foregroundEdgeCount " , request . getReferenceEdge ( ) . getForegroundEdge ( ) ! = null ?
request . getReferenceEdge ( ) . getForegroundEdge ( ) . length : 0 ) ;
response . put ( " backgroundEdgeCount " , request . getReferenceEdge ( ) . getBackgroundEdge ( ) ! = null ?
request . getReferenceEdge ( ) . getBackgroundEdge ( ) . length : 0 ) ;
}
log . info ( " 阿里云图像生成任务提交成功, 任务ID: {}, 处理时间: {}ms " , parsedResponse . get ( " taskId " ) , processingTime ) ;
} catch ( Exception e ) {
log . error ( " 阿里云图像生成失败: {} " , e . getMessage ( ) , e ) ;
response . put ( " success " , false ) ;
response . put ( " error " , e . getMessage ( ) ) ;
response . put ( " processingTimeMs " , System . currentTimeMillis ( ) - startTime ) ;
}
return response ;
}
/**
* 处理引导图像RGBA转换
*
* 下载引导图像, 转换为RGBA格式, 上传到MinIO并生成新的临时访问链接
*
* @param refImageUrl 原始引导图像URL
* @return 处理后的RGBA图像临时访问链接, 如果处理失败返回null
*/
private String processRefImageForRGBA ( String refImageUrl ) {
try {
log . info ( " 开始处理引导图像RGBA转换, 原始URL: {} " , refImageUrl ) ;
// 1. 下载原始图像
byte [ ] originalImageBytes = downloadImageFromUrl ( refImageUrl ) ;
log . info ( " 原始图像下载完成,大小: {} bytes " , originalImageBytes . length ) ;
// 2. 转换为RGBA格式
byte [ ] rgbaImageBytes = convertImageToRGBA ( originalImageBytes ) ;
log . info ( " 图像RGBA转换完成, 大小: {} bytes " , rgbaImageBytes . length ) ;
// 3. 生成唯一文件名
String uniqueFileName = generateUniqueFileName ( " png " ) ;
// 4. 上传RGBA图像到MinIO
String materialUrl = minIOService . uploadFile ( rgbaImageBytes , uniqueFileName , " image/png " ) ;
log . info ( " RGBA图像上传到MinIO成功: {} " , materialUrl ) ;
// 5. 生成7天临时访问链接
String materialTempUrl = minIOService . generateTempUrl ( uniqueFileName ) ;
log . info ( " 生成7天临时访问链接: {} " , materialTempUrl ) ;
return materialTempUrl ;
} catch ( Exception e ) {
log . error ( " 处理引导图像RGBA转换失败: {} " , refImageUrl , e ) ;
return null ;
}
}
/**
* 从URL下载图像
*
* @param imageUrl 图像URL
* @return 图像字节数组
* @throws IOException 下载失败时抛出异常
*/
private byte [ ] downloadImageFromUrl ( String imageUrl ) throws IOException {
try {
log . info ( " 开始下载图像: {} " , imageUrl ) ;
URL url = new URL ( imageUrl ) ;
try ( InputStream inputStream = url . openStream ( ) ;
ByteArrayOutputStream outputStream = new ByteArrayOutputStream ( ) ) {
byte [ ] buffer = new byte [ 4096 ] ;
int bytesRead ;
while ( ( bytesRead = inputStream . read ( buffer ) ) ! = - 1 ) {
outputStream . write ( buffer , 0 , bytesRead ) ;
}
byte [ ] imageBytes = outputStream . toByteArray ( ) ;
log . info ( " 图像下载完成,大小: {} bytes " , imageBytes . length ) ;
return imageBytes ;
}
} catch ( Exception e ) {
log . error ( " 下载图像失败: {} " , imageUrl , e ) ;
throw new IOException ( " 下载图像失败: " + e . getMessage ( ) , e ) ;
}
}
/**
* 将图像转换为RGBA格式
*
* @param imageBytes 原始图像字节数组
* @return RGBA格式的图像字节数组
* @throws IOException 转换失败时抛出异常
*/
private byte [ ] convertImageToRGBA ( byte [ ] imageBytes ) throws IOException {
try {
log . info ( " 开始转换图像为RGBA格式 " ) ;
// 从字节数组读取图像
ByteArrayInputStream bais = new ByteArrayInputStream ( imageBytes ) ;
java . awt . image . BufferedImage originalImage = javax . imageio . ImageIO . read ( bais ) ;
if ( originalImage = = null ) {
throw new IllegalArgumentException ( " 无法读取图像文件 " ) ;
}
int width = originalImage . getWidth ( ) ;
int height = originalImage . getHeight ( ) ;
// 验证图像尺寸
int maxDimension = Math . max ( width , height ) ;
if ( maxDimension > 2048 ) {
throw new IllegalArgumentException ( " 图像长边不能超过2048像素, 当前尺寸: " + width + " x " + height ) ;
}
log . info ( " 原始图像尺寸: {}x{} " , width , height ) ;
// 创建RGBA格式的图像
java . awt . image . BufferedImage rgbaImage = new java . awt . image . BufferedImage (
width , height , java . awt . image . BufferedImage . TYPE_INT_ARGB ) ;
// 获取图形上下文
java . awt . Graphics2D g2d = rgbaImage . createGraphics ( ) ;
// 设置渲染提示以获得更好的质量
g2d . setRenderingHint ( java . awt . RenderingHints . KEY_INTERPOLATION ,
java . awt . RenderingHints . VALUE_INTERPOLATION_BILINEAR ) ;
g2d . setRenderingHint ( java . awt . RenderingHints . KEY_RENDERING ,
java . awt . RenderingHints . VALUE_RENDER_QUALITY ) ;
g2d . setRenderingHint ( java . awt . RenderingHints . KEY_ANTIALIASING ,
java . awt . RenderingHints . VALUE_ANTIALIAS_ON ) ;
// 绘制原始图像到RGBA图像上
g2d . drawImage ( originalImage , 0 , 0 , null ) ;
g2d . dispose ( ) ;
// 将RGBA图像转换为字节数组
ByteArrayOutputStream baos = new ByteArrayOutputStream ( ) ;
javax . imageio . ImageIO . write ( rgbaImage , " PNG " , baos ) ;
byte [ ] rgbaBytes = baos . toByteArray ( ) ;
log . info ( " 图像RGBA转换完成, 大小: {} bytes " , rgbaBytes . length ) ;
return rgbaBytes ;
} catch ( Exception e ) {
log . error ( " 转换图像为RGBA格式失败 " , e ) ;
throw new IOException ( " 转换图像为RGBA格式失败: " + e . getMessage ( ) , e ) ;
}
}
/**
* 生成唯一文件名
*
* @param extension 文件扩展名
* @return 唯一文件名
*/
private String generateUniqueFileName ( String extension ) {
long timestamp = System . currentTimeMillis ( ) ;
String uuid = java . util . UUID . randomUUID ( ) . toString ( ) . replace ( " - " , " " ) ;
return timestamp + " _ " + uuid + " . " + extension ;
}
/**
* 构建阿里云请求参数
*/
private Map < String , Object > buildAliyunRequest ( TextModelController . ImageProcessingRequest request ) {
Map < String , Object > aliyunRequest = new HashMap < > ( ) ;
// 设置模型
aliyunRequest . put ( " model " , " wanx-background-generation-v2 " ) ;
// 构建input参数
Map < String , Object > input = new HashMap < > ( ) ;
input . put ( " base_image_url " , request . getBaseImageUrl ( ) ) ;
input . put ( " ref_image_url " , request . getRefImageUrl ( ) ) ;
input . put ( " ref_prompt " , request . getRefPrompt ( ) ) ;
// 构建reference_edge参数
if ( request . getReferenceEdge ( ) ! = null ) {
Map < String , Object > referenceEdge = new HashMap < > ( ) ;
referenceEdge . put ( " foreground_edge " , request . getReferenceEdge ( ) . getForegroundEdge ( ) ) ;
referenceEdge . put ( " background_edge " , request . getReferenceEdge ( ) . getBackgroundEdge ( ) ) ;
referenceEdge . put ( " foreground_edge_prompt " , request . getReferenceEdge ( ) . getForegroundEdgePrompt ( ) ) ;
referenceEdge . put ( " background_edge_prompt " , request . getReferenceEdge ( ) . getBackgroundEdgePrompt ( ) ) ;
input . put ( " reference_edge " , referenceEdge ) ;
}
aliyunRequest . put ( " input " , input ) ;
// 构建parameters参数
Map < String , Object > parameters = new HashMap < > ( ) ;
if ( request . getParameters ( ) ! = null ) {
parameters . put ( " n " , request . getParameters ( ) . getN ( ) ) ;
parameters . put ( " ref_prompt_weight " , request . getParameters ( ) . getRefPromptWeight ( ) ) ;
parameters . put ( " model_version " , request . getParameters ( ) . getModelVersion ( ) ) ;
} else {
parameters . put ( " n " , 1 ) ;
parameters . put ( " ref_prompt_weight " , 0 . 5 ) ;
parameters . put ( " model_version " , " v3 " ) ;
}
aliyunRequest . put ( " parameters " , parameters ) ;
return aliyunRequest ;
}
/**
* 发送请求到阿里云
*/
private String sendRequestToAliyun ( Map < String , Object > request ) throws Exception {
String apiKey = System . getenv ( " DASHSCOPE_API_KEY " ) ;
if ( apiKey = = null | | apiKey . isEmpty ( ) ) {
throw new RuntimeException ( " DASHSCOPE_API_KEY 环境变量未设置 " ) ;
}
String url = " https://dashscope.aliyuncs.com/api/v1/services/aigc/background-generation/generation/ " ;
HttpHeaders headers = new HttpHeaders ( ) ;
headers . setContentType ( MediaType . APPLICATION_JSON ) ;
headers . set ( " X-DashScope-Async " , " enable " ) ;
headers . set ( " Authorization " , " Bearer " + apiKey ) ;
HttpEntity < Map < String , Object > > entity = new HttpEntity < > ( request , headers ) ;
log . info ( " 发送请求到阿里云: {} " , objectMapper . writeValueAsString ( request ) ) ;
ResponseEntity < String > response = restTemplate . exchange ( url , HttpMethod . POST , entity , String . class ) ;
if ( response . getStatusCode ( ) . is2xxSuccessful ( ) ) {
log . info ( " 阿里云请求成功,响应: {} " , response . getBody ( ) ) ;
return response . getBody ( ) ;
} else {
throw new RuntimeException ( " 阿里云请求失败,状态码: " + response . getStatusCode ( ) ) ;
}
}
/**
* 解析阿里云响应
*/
private Map < String , Object > parseAliyunResponse ( String responseBody ) throws Exception {
JsonNode rootNode = objectMapper . readTree ( responseBody ) ;
Map < String , Object > result = new HashMap < > ( ) ;
if ( rootNode . has ( " output " ) ) {
JsonNode outputNode = rootNode . get ( " output " ) ;
if ( outputNode . has ( " task_id " ) ) {
result . put ( " taskId " , outputNode . get ( " task_id " ) . asText ( ) ) ;
}
if ( outputNode . has ( " task_status " ) ) {
result . put ( " status " , outputNode . get ( " task_status " ) . asText ( ) ) ;
}
}
if ( rootNode . has ( " request_id " ) ) {
result . put ( " requestId " , rootNode . get ( " request_id " ) . asText ( ) ) ;
}
return result ;
}
/**
* 创建ImageModel记录对象
*/
private ImageModel createImageModelRecord ( TextModelController . ImageProcessingRequest request , Map < String , Object > parsedResponse ) {
ImageModel imageModel = new ImageModel ( ) ;
// 设置基本信息
imageModel . setImageName ( request . getImageName ( ) ) ;
imageModel . setModelName ( " wanx-background-generation-v2 " ) ;
imageModel . setOwnerName ( request . getOwnerName ( ) ) ;
imageModel . setOwnerPhone ( request . getOwnerPhone ( ) ) ;
imageModel . setImageType ( " result " ) ; // 生成的结果图片
// 设置时间
imageModel . setCreateTime ( LocalDateTime . now ( ) ) ;
imageModel . setUpdateTime ( LocalDateTime . now ( ) ) ;
// 设置素材URL( 包含所有相关URL的JSON)
String materialUrlJson = createMaterialUrlJson ( request ) ;
imageModel . setMaterialTempUrl ( materialUrlJson ) ;
// 设置提示词
imageModel . setPrompt ( request . getRefPrompt ( ) ) ;
// 设置功能函数
imageModel . setFunctionName ( " background_generation " ) ;
// 设置请求ID、任务ID和任务状态
imageModel . setRequestId ( generateRequestId ( ) ) ;
imageModel . setTaskId ( ( String ) parsedResponse . get ( " taskId " ) ) ;
imageModel . setTaskStatus ( ( String ) parsedResponse . get ( " status " ) ) ;
log . info ( " 创建ImageModel记录: 所属人={}, 任务ID={}, 任务状态={}, 请求ID={} " ,
request . getOwnerName ( ) , parsedResponse . get ( " taskId " ) , parsedResponse . get ( " status " ) , imageModel . getRequestId ( ) ) ;
return imageModel ;
}
/**
* 创建包含所有URL信息的JSON字符串
*/
private String createMaterialUrlJson ( TextModelController . ImageProcessingRequest request ) {
try {
Map < String , Object > urlInfo = new HashMap < > ( ) ;
// 主图URL
urlInfo . put ( " baseImageUrl " , request . getBaseImageUrl ( ) ) ;
// 参考图URL
urlInfo . put ( " refImageUrl " , request . getRefImageUrl ( ) ) ;
// 参考边缘信息
if ( request . getReferenceEdge ( ) ! = null ) {
Map < String , Object > edgeInfo = new HashMap < > ( ) ;
// 前景边缘URL
if ( request . getReferenceEdge ( ) . getForegroundEdge ( ) ! = null ) {
edgeInfo . put ( " foregroundEdge " , request . getReferenceEdge ( ) . getForegroundEdge ( ) ) ;
}
// 背景边缘URL
if ( request . getReferenceEdge ( ) . getBackgroundEdge ( ) ! = null ) {
edgeInfo . put ( " backgroundEdge " , request . getReferenceEdge ( ) . getBackgroundEdge ( ) ) ;
}
// 前景边缘提示词
if ( request . getReferenceEdge ( ) . getForegroundEdgePrompt ( ) ! = null ) {
edgeInfo . put ( " foregroundEdgePrompt " , request . getReferenceEdge ( ) . getForegroundEdgePrompt ( ) ) ;
}
// 背景边缘提示词
if ( request . getReferenceEdge ( ) . getBackgroundEdgePrompt ( ) ! = null ) {
edgeInfo . put ( " backgroundEdgePrompt " , request . getReferenceEdge ( ) . getBackgroundEdgePrompt ( ) ) ;
}
urlInfo . put ( " referenceEdge " , edgeInfo ) ;
}
// 转换为JSON字符串
String jsonString = objectMapper . writeValueAsString ( urlInfo ) ;
log . info ( " 创建素材URL JSON: {} " , jsonString ) ;
return jsonString ;
} catch ( Exception e ) {
log . error ( " 创建素材URL JSON失败: {} " , e . getMessage ( ) , e ) ;
// 如果JSON创建失败, 至少保存基础图片URL
return request . getBaseImageUrl ( ) ;
}
}
/**
* 手动触发图像模型任务状态检查
*
* 提供手动触发任务状态检查的接口,用于:
* - 立即检查所有PENDING或RUNNING状态的任务
* - 调试和测试任务状态检查功能
* - 紧急情况下手动处理任务
*
* @return ResponseEntity 包含执行结果的响应实体
* - success: 执行是否成功
* - message: 执行结果消息
* - processedCount: 处理的任务数量
* - timestamp: 执行时间戳
*/
@PostMapping ( " /check-status " )
@Operation ( summary = " 手动触发任务状态检查 " , description = " 立即执行图像模型任务状态检查, 处理所有PENDING或RUNNING状态的任务 " )
public ResponseEntity < Map < String , Object > > checkImageModelStatus ( ) {
Map < String , Object > result = new HashMap < > ( ) ;
try {
log . info ( " 手动触发图像模型任务状态检查开始 " ) ;
// 调用调度器的检查方法
imageModelStatusScheduler . checkImageModelStatus ( ) ;
result . put ( " success " , true ) ;
result . put ( " message " , " 图像模型任务状态检查执行成功 " ) ;
result . put ( " timestamp " , LocalDateTime . now ( ) ) ;
result . put ( " requestId " , generateRequestId ( ) ) ;
log . info ( " 手动触发图像模型任务状态检查完成 " ) ;
return ResponseEntity . ok ( result ) ;
} catch ( Exception e ) {
log . error ( " 手动触发图像模型任务状态检查失败: {} " , e . getMessage ( ) , e ) ;
result . put ( " success " , false ) ;
result . put ( " message " , " 图像模型任务状态检查执行失败: " + e . getMessage ( ) ) ;
result . put ( " error " , e . getClass ( ) . getSimpleName ( ) ) ;
result . put ( " timestamp " , LocalDateTime . now ( ) ) ;
result . put ( " requestId " , generateRequestId ( ) ) ;
return ResponseEntity . status ( HttpStatus . INTERNAL_SERVER_ERROR ) . body ( result ) ;
}
}
/**
* 生成请求ID
*/
private String generateRequestId ( ) {
return " REQ_ " + System . currentTimeMillis ( ) + " _ " + ( int ) ( Math . random ( ) * 1000 ) ;
}
}