Commit 1a903560 by shun peng

fix: 修改对话名称自动填写、添加智能体接口

parent 576ea648
<?xml version="1.0" encoding="UTF-8"?>
<project version="4">
<component name="Palette2">
<group name="Swing">
<item class="com.intellij.uiDesigner.HSpacer" tooltip-text="Horizontal Spacer" icon="/com/intellij/uiDesigner/icons/hspacer.svg" removable="false" auto-create-binding="false" can-attach-label="false">
<default-constraints vsize-policy="1" hsize-policy="6" anchor="0" fill="1" />
</item>
<item class="com.intellij.uiDesigner.VSpacer" tooltip-text="Vertical Spacer" icon="/com/intellij/uiDesigner/icons/vspacer.svg" removable="false" auto-create-binding="false" can-attach-label="false">
<default-constraints vsize-policy="6" hsize-policy="1" anchor="0" fill="2" />
</item>
<item class="javax.swing.JPanel" icon="/com/intellij/uiDesigner/icons/panel.svg" removable="false" auto-create-binding="false" can-attach-label="false">
<default-constraints vsize-policy="3" hsize-policy="3" anchor="0" fill="3" />
</item>
<item class="javax.swing.JScrollPane" icon="/com/intellij/uiDesigner/icons/scrollPane.svg" removable="false" auto-create-binding="false" can-attach-label="true">
<default-constraints vsize-policy="7" hsize-policy="7" anchor="0" fill="3" />
</item>
<item class="javax.swing.JButton" icon="/com/intellij/uiDesigner/icons/button.svg" removable="false" auto-create-binding="true" can-attach-label="false">
<default-constraints vsize-policy="0" hsize-policy="3" anchor="0" fill="1" />
<initial-values>
<property name="text" value="Button" />
</initial-values>
</item>
<item class="javax.swing.JRadioButton" icon="/com/intellij/uiDesigner/icons/radioButton.svg" removable="false" auto-create-binding="true" can-attach-label="false">
<default-constraints vsize-policy="0" hsize-policy="3" anchor="8" fill="0" />
<initial-values>
<property name="text" value="RadioButton" />
</initial-values>
</item>
<item class="javax.swing.JCheckBox" icon="/com/intellij/uiDesigner/icons/checkBox.svg" removable="false" auto-create-binding="true" can-attach-label="false">
<default-constraints vsize-policy="0" hsize-policy="3" anchor="8" fill="0" />
<initial-values>
<property name="text" value="CheckBox" />
</initial-values>
</item>
<item class="javax.swing.JLabel" icon="/com/intellij/uiDesigner/icons/label.svg" removable="false" auto-create-binding="false" can-attach-label="false">
<default-constraints vsize-policy="0" hsize-policy="0" anchor="8" fill="0" />
<initial-values>
<property name="text" value="Label" />
</initial-values>
</item>
<item class="javax.swing.JTextField" icon="/com/intellij/uiDesigner/icons/textField.svg" removable="false" auto-create-binding="true" can-attach-label="true">
<default-constraints vsize-policy="0" hsize-policy="6" anchor="8" fill="1">
<preferred-size width="150" height="-1" />
</default-constraints>
</item>
<item class="javax.swing.JPasswordField" icon="/com/intellij/uiDesigner/icons/passwordField.svg" removable="false" auto-create-binding="true" can-attach-label="true">
<default-constraints vsize-policy="0" hsize-policy="6" anchor="8" fill="1">
<preferred-size width="150" height="-1" />
</default-constraints>
</item>
<item class="javax.swing.JFormattedTextField" icon="/com/intellij/uiDesigner/icons/formattedTextField.svg" removable="false" auto-create-binding="true" can-attach-label="true">
<default-constraints vsize-policy="0" hsize-policy="6" anchor="8" fill="1">
<preferred-size width="150" height="-1" />
</default-constraints>
</item>
<item class="javax.swing.JTextArea" icon="/com/intellij/uiDesigner/icons/textArea.svg" removable="false" auto-create-binding="true" can-attach-label="true">
<default-constraints vsize-policy="6" hsize-policy="6" anchor="0" fill="3">
<preferred-size width="150" height="50" />
</default-constraints>
</item>
<item class="javax.swing.JTextPane" icon="/com/intellij/uiDesigner/icons/textPane.svg" removable="false" auto-create-binding="true" can-attach-label="true">
<default-constraints vsize-policy="6" hsize-policy="6" anchor="0" fill="3">
<preferred-size width="150" height="50" />
</default-constraints>
</item>
<item class="javax.swing.JEditorPane" icon="/com/intellij/uiDesigner/icons/editorPane.svg" removable="false" auto-create-binding="true" can-attach-label="true">
<default-constraints vsize-policy="6" hsize-policy="6" anchor="0" fill="3">
<preferred-size width="150" height="50" />
</default-constraints>
</item>
<item class="javax.swing.JComboBox" icon="/com/intellij/uiDesigner/icons/comboBox.svg" removable="false" auto-create-binding="true" can-attach-label="true">
<default-constraints vsize-policy="0" hsize-policy="2" anchor="8" fill="1" />
</item>
<item class="javax.swing.JTable" icon="/com/intellij/uiDesigner/icons/table.svg" removable="false" auto-create-binding="true" can-attach-label="false">
<default-constraints vsize-policy="6" hsize-policy="6" anchor="0" fill="3">
<preferred-size width="150" height="50" />
</default-constraints>
</item>
<item class="javax.swing.JList" icon="/com/intellij/uiDesigner/icons/list.svg" removable="false" auto-create-binding="true" can-attach-label="false">
<default-constraints vsize-policy="6" hsize-policy="2" anchor="0" fill="3">
<preferred-size width="150" height="50" />
</default-constraints>
</item>
<item class="javax.swing.JTree" icon="/com/intellij/uiDesigner/icons/tree.svg" removable="false" auto-create-binding="true" can-attach-label="false">
<default-constraints vsize-policy="6" hsize-policy="6" anchor="0" fill="3">
<preferred-size width="150" height="50" />
</default-constraints>
</item>
<item class="javax.swing.JTabbedPane" icon="/com/intellij/uiDesigner/icons/tabbedPane.svg" removable="false" auto-create-binding="true" can-attach-label="false">
<default-constraints vsize-policy="3" hsize-policy="3" anchor="0" fill="3">
<preferred-size width="200" height="200" />
</default-constraints>
</item>
<item class="javax.swing.JSplitPane" icon="/com/intellij/uiDesigner/icons/splitPane.svg" removable="false" auto-create-binding="false" can-attach-label="false">
<default-constraints vsize-policy="3" hsize-policy="3" anchor="0" fill="3">
<preferred-size width="200" height="200" />
</default-constraints>
</item>
<item class="javax.swing.JSpinner" icon="/com/intellij/uiDesigner/icons/spinner.svg" removable="false" auto-create-binding="true" can-attach-label="true">
<default-constraints vsize-policy="0" hsize-policy="6" anchor="8" fill="1" />
</item>
<item class="javax.swing.JSlider" icon="/com/intellij/uiDesigner/icons/slider.svg" removable="false" auto-create-binding="true" can-attach-label="false">
<default-constraints vsize-policy="0" hsize-policy="6" anchor="8" fill="1" />
</item>
<item class="javax.swing.JSeparator" icon="/com/intellij/uiDesigner/icons/separator.svg" removable="false" auto-create-binding="false" can-attach-label="false">
<default-constraints vsize-policy="6" hsize-policy="6" anchor="0" fill="3" />
</item>
<item class="javax.swing.JProgressBar" icon="/com/intellij/uiDesigner/icons/progressbar.svg" removable="false" auto-create-binding="true" can-attach-label="false">
<default-constraints vsize-policy="0" hsize-policy="6" anchor="0" fill="1" />
</item>
<item class="javax.swing.JToolBar" icon="/com/intellij/uiDesigner/icons/toolbar.svg" removable="false" auto-create-binding="false" can-attach-label="false">
<default-constraints vsize-policy="0" hsize-policy="6" anchor="0" fill="1">
<preferred-size width="-1" height="20" />
</default-constraints>
</item>
<item class="javax.swing.JToolBar$Separator" icon="/com/intellij/uiDesigner/icons/toolbarSeparator.svg" removable="false" auto-create-binding="false" can-attach-label="false">
<default-constraints vsize-policy="0" hsize-policy="0" anchor="0" fill="1" />
</item>
<item class="javax.swing.JScrollBar" icon="/com/intellij/uiDesigner/icons/scrollbar.svg" removable="false" auto-create-binding="true" can-attach-label="false">
<default-constraints vsize-policy="6" hsize-policy="0" anchor="0" fill="2" />
</item>
</group>
</component>
</project>
\ No newline at end of file
package cn.iocoder.yudao.module.ai.config;
import org.springframework.context.annotation.Bean;
import org.springframework.context.annotation.Configuration;
import org.springframework.web.client.RestTemplate;
@Configuration
public class RestTemplateConfig {
@Bean
public RestTemplate restTemplate() {
return new RestTemplate();
}
}
package cn.iocoder.yudao.module.ai.controller.admin.agent;
import cn.iocoder.yudao.framework.common.pojo.CommonResult;
import cn.iocoder.yudao.framework.common.util.object.BeanUtils;
import cn.iocoder.yudao.module.ai.controller.admin.agent.vo.AgentCreateReqVO;
import cn.iocoder.yudao.module.ai.controller.admin.agent.vo.AgentMsgReqVO;
import cn.iocoder.yudao.module.ai.controller.admin.agent.vo.AgentUsedRespVO;
import cn.iocoder.yudao.module.ai.dal.dataobject.chat.AiChatConversationDO;
import cn.iocoder.yudao.module.ai.dal.dataobject.model.AiApiKeyDO;
import cn.iocoder.yudao.module.ai.dal.dataobject.model.AiChatModelDO;
import cn.iocoder.yudao.module.ai.dal.dataobject.model.AiChatRoleDO;
import cn.iocoder.yudao.module.ai.service.chat.AiChatConversationService;
import cn.iocoder.yudao.module.ai.service.model.AiApiKeyService;
import cn.iocoder.yudao.module.ai.service.model.AiChatModelService;
import cn.iocoder.yudao.module.ai.service.model.AiChatRoleService;
import io.swagger.v3.oas.annotations.Operation;
import io.swagger.v3.oas.annotations.tags.Tag;
import jakarta.annotation.Resource;
import jakarta.validation.Valid;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.http.*;
import org.springframework.validation.annotation.Validated;
import org.springframework.web.bind.annotation.*;
import org.springframework.web.client.RestTemplate;
import java.util.ArrayList;
import java.util.HashMap;
import java.util.List;
import java.util.Map;
import static cn.iocoder.yudao.framework.common.pojo.CommonResult.success;
import static cn.iocoder.yudao.framework.security.core.util.SecurityFrameworkUtils.getLoginUserId;
@Tag(name = "管理后台 - 智能体")
@RestController
@RequestMapping("/ai/agent")
@Validated
public class AgentChatController {
@Autowired
private RestTemplate restTemplate;
@Resource
private AiApiKeyService apiKeyService;
@Resource
private AiChatConversationService chatConversationService;
@Resource
private AiChatRoleService chatRoleService;
@Resource
private AiChatModelService chatModelService;
@PostMapping("/create-agent-records")
@Operation(summary = "创建对话过的智能体记录")
public CommonResult<Long> createChatAgentRecord(@RequestBody @Valid AgentCreateReqVO createReqVO) {
return success(chatConversationService.createUsedAgentRecord(createReqVO, getLoginUserId()));
}
@GetMapping("/get-used-agents")
@Operation(summary = "获得对话过的智能体列表")
public CommonResult<List<AgentUsedRespVO>> getChatAgentList() {
List<AiChatConversationDO> list = chatConversationService.getChatAgentListByUserId(getLoginUserId());
List<AgentUsedRespVO> result = new ArrayList<>(list.size());
for (int i = 0; i < list.size(); i++) {
AiChatConversationDO aiChatConversationDO = list.get(i);
AiChatRoleDO aiChatRoleDO = chatRoleService.getChatRole(aiChatConversationDO.getRoleId());
AgentUsedRespVO agentUsedRespVO = new AgentUsedRespVO();
agentUsedRespVO.setId(aiChatConversationDO.getId());
agentUsedRespVO.setRoleId(aiChatConversationDO.getRoleId());
agentUsedRespVO.setRoleAvatar(aiChatRoleDO.getAvatar());
agentUsedRespVO.setCreateTime(aiChatConversationDO.getCreateTime());
agentUsedRespVO.setRoleName(aiChatRoleDO.getName());
result.add(agentUsedRespVO);
}
return success(result);
}
@PostMapping("/get-messages")
@Operation(summary = "获得智能体对话消息")
public ResponseEntity<?> getAgentMessages(@RequestBody @Valid AgentMsgReqVO reqVO) {
AiChatRoleDO aiRole = chatRoleService.getChatRole(reqVO.getRoleId());
AiChatConversationDO conversation = chatConversationService.getChatConversation(reqVO.getConversationId());
AiChatModelDO model = chatModelService.getChatModel(conversation.getModelId());
AiApiKeyDO apiKeyDO = apiKeyService.getApiKey(model.getKeyId());
// 目标服务的URL
String baseUrl = apiKeyDO.getUrl();
// String baseUrl = "http://218.77.58.8:8088/api";
// 构建请求参数
String appId = aiRole.getAppId();
String chatId = String.valueOf(reqVO.getConversationId());
// String appId = "66f225d52b2887652b418f82";
// String chatId = "1781604279872581766";
boolean loadCustomFeedbacks = true;
// 构建URL与请求参数
String requestUrl = baseUrl + "/core/chat/init?appId=" + appId + "&chatId=" + chatId + "&loadCustomFeedbacks=" + loadCustomFeedbacks;
// 设置请求头
HttpHeaders headers = new HttpHeaders();
headers.setContentType(MediaType.APPLICATION_JSON);
headers.set("Authorization", "Bearer " + apiKeyDO.getApiKey());
// 构建请求实体
HttpEntity<String> entity = new HttpEntity<>(headers);
// 发送请求并接收响应
ResponseEntity<String> response = restTemplate.exchange(requestUrl, HttpMethod.GET, entity, String.class);
// 返回响应
return response;
}
// @PostMapping("/list")
// @Operation(summary = "获得智能体列表")
// public ResponseEntity<?> getAgentList(@RequestBody @Valid AgentMsgReqVO reqVO) {
// // 构建请求参数
// Map<String, String> params = new HashMap<>();
// params.put("parentId", reqVO.getParentId());
// params.put("searchKey", reqVO.getSearchKey());
//
//
// AiApiKeyDO apiKeyDO = apiKeyService.getApiKeyByName(reqVO.getName());
//
// // 目标服务的URL
// String url = apiKeyDO.getUrl();
//
// // 设置请求头
// HttpHeaders headers = new HttpHeaders();
// headers.setContentType(MediaType.APPLICATION_JSON);
// headers.set("Authorization", "Bearer " + apiKeyDO.getApiKey());
//
// // 构建请求实体
// HttpEntity<Map<String, String>> entity = new HttpEntity<>(params, headers);
//
// // 发送请求并接收响应
// ResponseEntity<String> response = restTemplate.exchange(url, HttpMethod.GET, entity, String.class);
//
// // 返回响应
// return response;
// }
}
package cn.iocoder.yudao.module.ai.controller.admin.agent.vo;
import io.swagger.v3.oas.annotations.media.Schema;
import lombok.Data;
@Schema(description = "管理后台 - Agent Record Create Request VO")
@Data
public class AgentCreateReqVO {
@Schema(description = "智能体编号", requiredMode = Schema.RequiredMode.REQUIRED, example = "10")
private Long roleId;
}
package cn.iocoder.yudao.module.ai.controller.admin.agent.vo;
import io.swagger.v3.oas.annotations.media.Schema;
import lombok.Data;
@Schema(description = "管理后台 - Agent Message Request VO")
@Data
public class AgentMsgReqVO {
@Schema(description = "智能体编号", requiredMode = Schema.RequiredMode.REQUIRED, example = "10")
private Long roleId;
@Schema(description = "对话编号", requiredMode = Schema.RequiredMode.REQUIRED, example = "1024")
private Long conversationId;
}
package cn.iocoder.yudao.module.ai.controller.admin.agent.vo;
import cn.iocoder.yudao.module.ai.dal.dataobject.model.AiChatModelDO;
import cn.iocoder.yudao.module.ai.dal.dataobject.model.AiChatRoleDO;
import com.fhs.core.trans.anno.Trans;
import com.fhs.core.trans.constant.TransType;
import com.fhs.core.trans.vo.VO;
import io.swagger.v3.oas.annotations.media.Schema;
import lombok.Data;
import java.time.LocalDateTime;
@Schema(description = "管理后台 - Used Agent List Request VO")
@Data
public class AgentUsedRespVO {
@Schema(description = "对话编号", requiredMode = Schema.RequiredMode.REQUIRED, example = "1024")
private Long id;
@Schema(description = "智能体编号", example = "1")
@Trans(type = TransType.SIMPLE, target = AiChatRoleDO.class, fields = {"name", "avatar"}, refs = {"roleName", "roleAvatar"})
private Long roleId;
@Schema(description = "创建时间", requiredMode = Schema.RequiredMode.REQUIRED)
private LocalDateTime createTime;
// ========== 关联 role 信息 ==========
@Schema(description = "智能体头像", example = "https://www.iocoder.cn/1.png")
private String roleAvatar;
@Schema(description = "智能体名字", example = "小黄")
private String roleName;
}
...@@ -2,6 +2,8 @@ package cn.iocoder.yudao.module.ai.controller.admin.chat; ...@@ -2,6 +2,8 @@ package cn.iocoder.yudao.module.ai.controller.admin.chat;
import cn.hutool.core.collection.CollUtil; import cn.hutool.core.collection.CollUtil;
import cn.hutool.core.util.ObjUtil; import cn.hutool.core.util.ObjUtil;
import cn.iocoder.yudao.framework.ai.core.enums.AiPlatformEnum;
import cn.iocoder.yudao.framework.ai.core.factory.AiModelFactory;
import cn.iocoder.yudao.framework.common.pojo.CommonResult; import cn.iocoder.yudao.framework.common.pojo.CommonResult;
import cn.iocoder.yudao.framework.common.pojo.PageResult; import cn.iocoder.yudao.framework.common.pojo.PageResult;
import cn.iocoder.yudao.framework.common.util.collection.MapUtils; import cn.iocoder.yudao.framework.common.util.collection.MapUtils;
...@@ -23,6 +25,13 @@ import jakarta.annotation.Resource; ...@@ -23,6 +25,13 @@ import jakarta.annotation.Resource;
import jakarta.annotation.security.PermitAll; import jakarta.annotation.security.PermitAll;
import jakarta.validation.Valid; import jakarta.validation.Valid;
import lombok.extern.slf4j.Slf4j; import lombok.extern.slf4j.Slf4j;
import org.springframework.ai.chat.client.ChatClient;
import org.springframework.ai.chat.messages.AssistantMessage;
import org.springframework.ai.chat.messages.UserMessage;
import org.springframework.ai.chat.model.ChatModel;
import org.springframework.ai.chat.prompt.ChatOptions;
import org.springframework.ai.chat.prompt.Prompt;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.http.MediaType; import org.springframework.http.MediaType;
import org.springframework.security.access.prepost.PreAuthorize; import org.springframework.security.access.prepost.PreAuthorize;
import org.springframework.web.bind.annotation.*; import org.springframework.web.bind.annotation.*;
...@@ -56,10 +65,24 @@ public class AiChatMessageController { ...@@ -56,10 +65,24 @@ public class AiChatMessageController {
return success(chatMessageService.sendMessage(sendReqVO, getLoginUserId())); return success(chatMessageService.sendMessage(sendReqVO, getLoginUserId()));
} }
@Resource
private AiModelFactory modelFactory;
@Operation(summary = "test", description = "test")
@PermitAll
@PostMapping("/test")
public String test() {
ChatModel model = modelFactory.getOrCreateChatModel(AiPlatformEnum.OPENAI,
"fastgpt-x2RF9hkc90AUk3iFSsNaWphpA1jV78i9kUTNsmogjqpKQJcEamEvIrpq3T", "http://218.77.58.8:8088/api");
return model.call(new UserMessage("你是谁"), new AssistantMessage("我是孙悟空"), new UserMessage("用python写helloworld代码"));
}
@Operation(summary = "发送消息(流式)", description = "流式返回,响应较快") @Operation(summary = "发送消息(流式)", description = "流式返回,响应较快")
@PostMapping(value = "/send-stream", produces = MediaType.TEXT_EVENT_STREAM_VALUE) @PostMapping(value = "/send-stream", produces = MediaType.TEXT_EVENT_STREAM_VALUE)
@PermitAll // 解决 SSE 最终响应的时候,会被 Access Denied 拦截的问题 @PermitAll // 解决 SSE 最终响应的时候,会被 Access Denied 拦截的问题
public Flux<CommonResult<AiChatMessageSendRespVO>> sendChatMessageStream(@Valid @RequestBody AiChatMessageSendReqVO sendReqVO) throws MalformedURLException { public Flux<CommonResult<AiChatMessageSendRespVO>> sendChatMessageStream(@Valid @RequestBody AiChatMessageSendReqVO sendReqVO) throws MalformedURLException {
// 禁用ruoyi的上下文能力
sendReqVO.setUseContext(false);
return chatMessageService.sendChatMessageStream(sendReqVO, getLoginUserId()); return chatMessageService.sendChatMessageStream(sendReqVO, getLoginUserId());
} }
......
...@@ -54,4 +54,7 @@ public class AiChatRoleRespVO implements VO { ...@@ -54,4 +54,7 @@ public class AiChatRoleRespVO implements VO {
@Schema(description = "创建时间", requiredMode = Schema.RequiredMode.REQUIRED) @Schema(description = "创建时间", requiredMode = Schema.RequiredMode.REQUIRED)
private LocalDateTime createTime; private LocalDateTime createTime;
@Schema(description = "appId for fastgpt")
private String appId;
} }
\ No newline at end of file
...@@ -29,4 +29,7 @@ public class AiChatRoleSaveMyReqVO { ...@@ -29,4 +29,7 @@ public class AiChatRoleSaveMyReqVO {
@NotEmpty(message = "角色设定不能为空") @NotEmpty(message = "角色设定不能为空")
private String systemMessage; private String systemMessage;
@Schema(description = "appId for fastgpt", example = "66f225d52b2887652b418f82")
private String appId;
} }
\ No newline at end of file
...@@ -51,4 +51,7 @@ public class AiChatRoleSaveReqVO { ...@@ -51,4 +51,7 @@ public class AiChatRoleSaveReqVO {
@InEnum(CommonStatusEnum.class) @InEnum(CommonStatusEnum.class)
private Integer status; private Integer status;
@Schema(description = "appId for fastgpt", example = "66f225d52b2887652b418f82")
private String appId;
} }
\ No newline at end of file
...@@ -79,4 +79,6 @@ public class AiChatRoleDO extends BaseDO { ...@@ -79,4 +79,6 @@ public class AiChatRoleDO extends BaseDO {
*/ */
private Integer status; private Integer status;
private String appId;
} }
...@@ -35,4 +35,14 @@ public interface AiChatConversationMapper extends BaseMapperX<AiChatConversation ...@@ -35,4 +35,14 @@ public interface AiChatConversationMapper extends BaseMapperX<AiChatConversation
.orderByDesc(AiChatConversationDO::getId)); .orderByDesc(AiChatConversationDO::getId));
} }
default List<AiChatConversationDO> selectListWithCon(Long userId){
return selectList(new LambdaQueryWrapperX<AiChatConversationDO>()
.eq(AiChatConversationDO::getUserId, userId)
.isNotNull(AiChatConversationDO::getRoleId));
}
default AiChatConversationDO selectByRoleId(Long roleId) {
return selectOne(new LambdaQueryWrapperX<AiChatConversationDO>()
.eq(AiChatConversationDO::getRoleId, roleId));
}
} }
...@@ -56,4 +56,10 @@ public interface AiChatMessageMapper extends BaseMapperX<AiChatMessageDO> { ...@@ -56,4 +56,10 @@ public interface AiChatMessageMapper extends BaseMapperX<AiChatMessageDO> {
.orderByDesc(AiChatMessageDO::getId)); .orderByDesc(AiChatMessageDO::getId));
} }
default Long selectCounts(Long conId){
return selectCount(new LambdaQueryWrapperX<AiChatMessageDO>()
.eqIfPresent(AiChatMessageDO::getConversationId, conId)
.eqIfPresent(AiChatMessageDO::getDeleted, 0)
);
}
} }
package cn.iocoder.yudao.module.ai.service.chat; package cn.iocoder.yudao.module.ai.service.chat;
import cn.iocoder.yudao.framework.common.pojo.PageResult; import cn.iocoder.yudao.framework.common.pojo.PageResult;
import cn.iocoder.yudao.module.ai.controller.admin.agent.vo.AgentCreateReqVO;
import cn.iocoder.yudao.module.ai.controller.admin.chat.vo.conversation.AiChatConversationCreateMyReqVO; import cn.iocoder.yudao.module.ai.controller.admin.chat.vo.conversation.AiChatConversationCreateMyReqVO;
import cn.iocoder.yudao.module.ai.controller.admin.chat.vo.conversation.AiChatConversationPageReqVO; import cn.iocoder.yudao.module.ai.controller.admin.chat.vo.conversation.AiChatConversationPageReqVO;
import cn.iocoder.yudao.module.ai.controller.admin.chat.vo.conversation.AiChatConversationUpdateMyReqVO; import cn.iocoder.yudao.module.ai.controller.admin.chat.vo.conversation.AiChatConversationUpdateMyReqVO;
...@@ -87,4 +88,7 @@ public interface AiChatConversationService { ...@@ -87,4 +88,7 @@ public interface AiChatConversationService {
*/ */
PageResult<AiChatConversationDO> getChatConversationPage(AiChatConversationPageReqVO pageReqVO); PageResult<AiChatConversationDO> getChatConversationPage(AiChatConversationPageReqVO pageReqVO);
List<AiChatConversationDO> getChatAgentListByUserId(Long loginUserId);
Long createUsedAgentRecord(AgentCreateReqVO createReqVO, Long loginUserId);
} }
...@@ -6,6 +6,7 @@ import cn.hutool.core.util.ObjUtil; ...@@ -6,6 +6,7 @@ import cn.hutool.core.util.ObjUtil;
import cn.hutool.core.util.ObjectUtil; import cn.hutool.core.util.ObjectUtil;
import cn.iocoder.yudao.framework.common.pojo.PageResult; import cn.iocoder.yudao.framework.common.pojo.PageResult;
import cn.iocoder.yudao.framework.common.util.object.BeanUtils; import cn.iocoder.yudao.framework.common.util.object.BeanUtils;
import cn.iocoder.yudao.module.ai.controller.admin.agent.vo.AgentCreateReqVO;
import cn.iocoder.yudao.module.ai.controller.admin.chat.vo.conversation.AiChatConversationCreateMyReqVO; import cn.iocoder.yudao.module.ai.controller.admin.chat.vo.conversation.AiChatConversationCreateMyReqVO;
import cn.iocoder.yudao.module.ai.controller.admin.chat.vo.conversation.AiChatConversationPageReqVO; import cn.iocoder.yudao.module.ai.controller.admin.chat.vo.conversation.AiChatConversationPageReqVO;
import cn.iocoder.yudao.module.ai.controller.admin.chat.vo.conversation.AiChatConversationUpdateMyReqVO; import cn.iocoder.yudao.module.ai.controller.admin.chat.vo.conversation.AiChatConversationUpdateMyReqVO;
...@@ -154,4 +155,36 @@ public class AiChatConversationServiceImpl implements AiChatConversationService ...@@ -154,4 +155,36 @@ public class AiChatConversationServiceImpl implements AiChatConversationService
return chatConversationMapper.selectChatConversationPage(pageReqVO); return chatConversationMapper.selectChatConversationPage(pageReqVO);
} }
@Override
public List<AiChatConversationDO> getChatAgentListByUserId(Long userId){
return chatConversationMapper.selectListWithCon(userId);
}
@Override
public Long createUsedAgentRecord(AgentCreateReqVO createReqVO, Long userId) {
// 1.1 获得 AiChatRoleDO 聊天角色
AiChatRoleDO role = chatRoleService.validateChatRole(createReqVO.getRoleId());
// 1.2 获得 AiChatModelDO 聊天模型
AiChatModelDO model = chatModalService.validateChatModel(role.getModelId());
Assert.notNull(model, "智能体未配置推理模型");
validateChatModel(model);
AiChatConversationDO conversationExist= chatConversationMapper.selectByRoleId(createReqVO.getRoleId());
if(conversationExist != null){
return conversationExist.getId();
}
// 2. 创建 AiChatConversationDO 聊天对话
AiChatConversationDO conversation = new AiChatConversationDO().setUserId(userId).setPinned(false)
.setModelId(model.getId()).setModel(model.getModel())
.setTemperature(model.getTemperature()).setMaxTokens(model.getMaxTokens()).setMaxContexts(model.getMaxContexts());
if (role != null) {
conversation.setTitle(role.getName()).setRoleId(role.getId()).setSystemMessage(role.getSystemMessage());
} else {
conversation.setTitle(AiChatConversationDO.TITLE_DEFAULT);
}
chatConversationMapper.insert(conversation);
return conversation.getId();
}
} }
...@@ -9,6 +9,8 @@ import cn.iocoder.yudao.framework.common.pojo.CommonResult; ...@@ -9,6 +9,8 @@ import cn.iocoder.yudao.framework.common.pojo.CommonResult;
import cn.iocoder.yudao.framework.common.pojo.PageResult; import cn.iocoder.yudao.framework.common.pojo.PageResult;
import cn.iocoder.yudao.framework.common.util.object.BeanUtils; import cn.iocoder.yudao.framework.common.util.object.BeanUtils;
import cn.iocoder.yudao.framework.tenant.core.util.TenantUtils; import cn.iocoder.yudao.framework.tenant.core.util.TenantUtils;
import cn.iocoder.yudao.module.ai.controller.admin.chat.vo.conversation.AiChatConversationRespVO;
import cn.iocoder.yudao.module.ai.controller.admin.chat.vo.conversation.AiChatConversationUpdateMyReqVO;
import cn.iocoder.yudao.module.ai.controller.admin.chat.vo.message.AiChatMessagePageReqVO; import cn.iocoder.yudao.module.ai.controller.admin.chat.vo.message.AiChatMessagePageReqVO;
import cn.iocoder.yudao.module.ai.controller.admin.chat.vo.message.AiChatMessageSendReqVO; import cn.iocoder.yudao.module.ai.controller.admin.chat.vo.message.AiChatMessageSendReqVO;
import cn.iocoder.yudao.module.ai.controller.admin.chat.vo.message.AiChatMessageSendRespVO; import cn.iocoder.yudao.module.ai.controller.admin.chat.vo.message.AiChatMessageSendRespVO;
...@@ -107,12 +109,22 @@ public class AiChatMessageServiceImpl implements AiChatMessageService { ...@@ -107,12 +109,22 @@ public class AiChatMessageServiceImpl implements AiChatMessageService {
StreamingChatModel chatModel = apiKeyService.getChatModel(model.getKeyId()); StreamingChatModel chatModel = apiKeyService.getChatModel(model.getKeyId());
// 2. 插入 user 发送消息 // 2. 插入 user 发送消息
Long count = chatMessageMapper.selectCounts(conversation.getId());
if(count == 0){
AiChatConversationUpdateMyReqVO updateMyReq = BeanUtils.toBean(conversation, AiChatConversationUpdateMyReqVO.class);
updateMyReq.setTitle(sendReqVO.getContent());
chatConversationService.updateChatConversationMy(updateMyReq, userId);
}
AiChatMessageDO userMessage = createChatMessage(conversation.getId(), null, model, AiChatMessageDO userMessage = createChatMessage(conversation.getId(), null, model,
userId, conversation.getRoleId(), MessageType.USER, sendReqVO.getContent(), sendReqVO.getUseContext()); userId, conversation.getRoleId(), MessageType.USER, sendReqVO.getContent(), sendReqVO.getUseContext());
// 3.1 插入 assistant 接收消息 // 3.1 插入 assistant 接收消息
AiChatMessageDO assistantMessage = createChatMessage(conversation.getId(), userMessage.getId(), model, // AiChatMessageDO assistantMessage = createChatMessage(conversation.getId(), userMessage.getId(), model,
userId, conversation.getRoleId(), MessageType.ASSISTANT, "", sendReqVO.getUseContext()); // userId, conversation.getRoleId(), MessageType.ASSISTANT, "", sendReqVO.getUseContext());
AiChatMessageDO assistantMessage = new AiChatMessageDO().setConversationId(conversation.getId()).setReplyId(userMessage.getId())
.setModel(model.getModel()).setModelId(model.getId()).setUserId(userId).setRoleId(conversation.getRoleId())
.setType(MessageType.ASSISTANT.getValue()).setContent("").setUseContext(sendReqVO.getUseContext());
assistantMessage.setCreateTime(LocalDateTime.now());
// 3.2 构建 Prompt,并进行调用 // 3.2 构建 Prompt,并进行调用
Prompt prompt = buildPrompt(conversation, historyMessages, model, sendReqVO); Prompt prompt = buildPrompt(conversation, historyMessages, model, sendReqVO);
...@@ -130,13 +142,20 @@ public class AiChatMessageServiceImpl implements AiChatMessageService { ...@@ -130,13 +142,20 @@ public class AiChatMessageServiceImpl implements AiChatMessageService {
.setReceive(BeanUtils.toBean(assistantMessage, AiChatMessageSendRespVO.Message.class).setContent(newContent))); .setReceive(BeanUtils.toBean(assistantMessage, AiChatMessageSendRespVO.Message.class).setContent(newContent)));
}).doOnComplete(() -> { }).doOnComplete(() -> {
// 忽略租户,因为 Flux 异步无法透传租户 // 忽略租户,因为 Flux 异步无法透传租户
TenantUtils.executeIgnore(() -> // TenantUtils.executeIgnore(() -> chatMessageMapper.updateById(new AiChatMessageDO().setId(assistantMessage.getId()).setContent(contentBuffer.toString())));
chatMessageMapper.updateById(new AiChatMessageDO().setId(assistantMessage.getId()).setContent(contentBuffer.toString()))); TenantUtils.executeIgnore(() -> {
assistantMessage.setContent(contentBuffer.toString());
chatMessageMapper.insert(assistantMessage);
});
}).doOnError(throwable -> { }).doOnError(throwable -> {
log.error("[sendChatMessageStream][userId({}) sendReqVO({}) 发生异常]", userId, sendReqVO, throwable); log.error("[sendChatMessageStream][userId({}) sendReqVO({}) 发生异常]", userId, sendReqVO, throwable);
// 忽略租户,因为 Flux 异步无法透传租户 // 忽略租户,因为 Flux 异步无法透传租户
TenantUtils.executeIgnore(() -> // TenantUtils.executeIgnore(() ->
chatMessageMapper.updateById(new AiChatMessageDO().setId(assistantMessage.getId()).setContent(throwable.getMessage()))); // chatMessageMapper.updateById(new AiChatMessageDO().setId(assistantMessage.getId()).setContent(throwable.getMessage())));
TenantUtils.executeIgnore(() -> {
assistantMessage.setContent(throwable.getMessage());
chatMessageMapper.insert(assistantMessage);
});
}).onErrorResume(error -> Flux.just(error(ErrorCodeConstants.CHAT_STREAM_ERROR))); }).onErrorResume(error -> Flux.just(error(ErrorCodeConstants.CHAT_STREAM_ERROR)));
} }
...@@ -164,7 +183,7 @@ public class AiChatMessageServiceImpl implements AiChatMessageService { ...@@ -164,7 +183,7 @@ public class AiChatMessageServiceImpl implements AiChatMessageService {
// 2. 构建 ChatOptions 对象 // 2. 构建 ChatOptions 对象
AiPlatformEnum platform = AiPlatformEnum.validatePlatform(model.getPlatform()); AiPlatformEnum platform = AiPlatformEnum.validatePlatform(model.getPlatform());
ChatOptions chatOptions = AiUtils.buildChatOptions(platform, model.getModel(), ChatOptions chatOptions = AiUtils.buildChatOptions(platform, model.getModel(),
conversation.getTemperature(), conversation.getMaxTokens()); conversation.getTemperature(), conversation.getMaxTokens(), String.valueOf(conversation.getId()));
return new Prompt(chatMessages, chatOptions); return new Prompt(chatMessages, chatOptions);
} }
......
...@@ -108,7 +108,7 @@ public class AiMindMapServiceImpl implements AiMindMapService { ...@@ -108,7 +108,7 @@ public class AiMindMapServiceImpl implements AiMindMapService {
List<Message> chatMessages = buildMessages(generateReqVO, systemMessage); List<Message> chatMessages = buildMessages(generateReqVO, systemMessage);
// 2. 构建 options 对象 // 2. 构建 options 对象
AiPlatformEnum platform = AiPlatformEnum.validatePlatform(model.getPlatform()); AiPlatformEnum platform = AiPlatformEnum.validatePlatform(model.getPlatform());
ChatOptions options = AiUtils.buildChatOptions(platform, model.getModel(), model.getTemperature(), model.getMaxTokens()); ChatOptions options = AiUtils.buildChatOptions(platform, model.getModel(), model.getTemperature(), model.getMaxTokens(), "");
return new Prompt(chatMessages, options); return new Prompt(chatMessages, options);
} }
......
...@@ -53,6 +53,14 @@ public interface AiApiKeyService { ...@@ -53,6 +53,14 @@ public interface AiApiKeyService {
AiApiKeyDO getApiKey(Long id); AiApiKeyDO getApiKey(Long id);
/** /**
* 通过名称获得 API 密钥
*
* @param name 名称
* @return API 密钥
*/
AiApiKeyDO getApiKeyByName(String name);
/**
* 校验 API 密钥 * 校验 API 密钥
* *
* @param id 比那好 * @param id 比那好
......
...@@ -12,6 +12,7 @@ import cn.iocoder.yudao.module.ai.controller.admin.model.vo.apikey.AiApiKeyPageR ...@@ -12,6 +12,7 @@ import cn.iocoder.yudao.module.ai.controller.admin.model.vo.apikey.AiApiKeyPageR
import cn.iocoder.yudao.module.ai.controller.admin.model.vo.apikey.AiApiKeySaveReqVO; import cn.iocoder.yudao.module.ai.controller.admin.model.vo.apikey.AiApiKeySaveReqVO;
import cn.iocoder.yudao.module.ai.dal.dataobject.model.AiApiKeyDO; import cn.iocoder.yudao.module.ai.dal.dataobject.model.AiApiKeyDO;
import cn.iocoder.yudao.module.ai.dal.mysql.model.AiApiKeyMapper; import cn.iocoder.yudao.module.ai.dal.mysql.model.AiApiKeyMapper;
import com.baomidou.mybatisplus.core.conditions.query.QueryWrapper;
import jakarta.annotation.Resource; import jakarta.annotation.Resource;
import org.springframework.ai.chat.model.ChatModel; import org.springframework.ai.chat.model.ChatModel;
import org.springframework.ai.embedding.EmbeddingModel; import org.springframework.ai.embedding.EmbeddingModel;
...@@ -80,6 +81,12 @@ public class AiApiKeyServiceImpl implements AiApiKeyService { ...@@ -80,6 +81,12 @@ public class AiApiKeyServiceImpl implements AiApiKeyService {
public AiApiKeyDO getApiKey(Long id) { public AiApiKeyDO getApiKey(Long id) {
return apiKeyMapper.selectById(id); return apiKeyMapper.selectById(id);
} }
@Override
public AiApiKeyDO getApiKeyByName(String name) {
QueryWrapper<AiApiKeyDO> queryWrapper = new QueryWrapper<>();
queryWrapper.eq("name", name);
return apiKeyMapper.selectOne(queryWrapper);
}
@Override @Override
public AiApiKeyDO validateApiKey(Long id) { public AiApiKeyDO validateApiKey(Long id) {
......
...@@ -126,7 +126,7 @@ public class AiWriteServiceImpl implements AiWriteService { ...@@ -126,7 +126,7 @@ public class AiWriteServiceImpl implements AiWriteService {
List<Message> chatMessages = buildMessages(generateReqVO, systemMessage); List<Message> chatMessages = buildMessages(generateReqVO, systemMessage);
// 2. 构建 options 对象 // 2. 构建 options 对象
AiPlatformEnum platform = AiPlatformEnum.validatePlatform(model.getPlatform()); AiPlatformEnum platform = AiPlatformEnum.validatePlatform(model.getPlatform());
ChatOptions options = AiUtils.buildChatOptions(platform, model.getModel(), model.getTemperature(), model.getMaxTokens()); ChatOptions options = AiUtils.buildChatOptions(platform, model.getModel(), model.getTemperature(), model.getMaxTokens(), "");
return new Prompt(chatMessages, options); return new Prompt(chatMessages, options);
} }
......
...@@ -30,7 +30,7 @@ public enum AiPlatformEnum { ...@@ -30,7 +30,7 @@ public enum AiPlatformEnum {
MIDJOURNEY("Midjourney", "Midjourney"), // Midjourney MIDJOURNEY("Midjourney", "Midjourney"), // Midjourney
SUNO("Suno", "Suno"), // Suno AI SUNO("Suno", "Suno"), // Suno AI
; FAST_GPT("FastGPT", "FastGPT");
/** /**
* 平台 * 平台
......
...@@ -10,6 +10,8 @@ import cn.iocoder.yudao.framework.ai.config.YudaoAiAutoConfiguration; ...@@ -10,6 +10,8 @@ import cn.iocoder.yudao.framework.ai.config.YudaoAiAutoConfiguration;
import cn.iocoder.yudao.framework.ai.config.YudaoAiProperties; import cn.iocoder.yudao.framework.ai.config.YudaoAiProperties;
import cn.iocoder.yudao.framework.ai.core.enums.AiPlatformEnum; import cn.iocoder.yudao.framework.ai.core.enums.AiPlatformEnum;
import cn.iocoder.yudao.framework.ai.core.model.deepseek.DeepSeekChatModel; import cn.iocoder.yudao.framework.ai.core.model.deepseek.DeepSeekChatModel;
import cn.iocoder.yudao.framework.ai.core.model.fastgpt.FastGPTApi;
import cn.iocoder.yudao.framework.ai.core.model.fastgpt.FastGPTChatModel;
import cn.iocoder.yudao.framework.ai.core.model.midjourney.api.MidjourneyApi; import cn.iocoder.yudao.framework.ai.core.model.midjourney.api.MidjourneyApi;
import cn.iocoder.yudao.framework.ai.core.model.suno.api.SunoApi; import cn.iocoder.yudao.framework.ai.core.model.suno.api.SunoApi;
import cn.iocoder.yudao.framework.ai.core.model.xinghuo.XingHuoChatModel; import cn.iocoder.yudao.framework.ai.core.model.xinghuo.XingHuoChatModel;
...@@ -93,6 +95,8 @@ public class AiModelFactoryImpl implements AiModelFactory { ...@@ -93,6 +95,8 @@ public class AiModelFactoryImpl implements AiModelFactory {
return buildAzureOpenAiChatModel(apiKey, url); return buildAzureOpenAiChatModel(apiKey, url);
case OLLAMA: case OLLAMA:
return buildOllamaChatModel(url); return buildOllamaChatModel(url);
case FAST_GPT:
return buildFastGPTModel(apiKey, url);
default: default:
throw new IllegalArgumentException(StrUtil.format("未知平台({})", platform)); throw new IllegalArgumentException(StrUtil.format("未知平台({})", platform));
} }
...@@ -119,6 +123,8 @@ public class AiModelFactoryImpl implements AiModelFactory { ...@@ -119,6 +123,8 @@ public class AiModelFactoryImpl implements AiModelFactory {
return SpringUtil.getBean(AzureOpenAiChatModel.class); return SpringUtil.getBean(AzureOpenAiChatModel.class);
case OLLAMA: case OLLAMA:
return SpringUtil.getBean(OllamaChatModel.class); return SpringUtil.getBean(OllamaChatModel.class);
case FAST_GPT:
return SpringUtil.getBean(FastGPTChatModel.class);
default: default:
throw new IllegalArgumentException(StrUtil.format("未知平台({})", platform)); throw new IllegalArgumentException(StrUtil.format("未知平台({})", platform));
} }
...@@ -340,4 +346,10 @@ public class AiModelFactoryImpl implements AiModelFactory { ...@@ -340,4 +346,10 @@ public class AiModelFactoryImpl implements AiModelFactory {
return new TongYiAutoConfiguration().tongYiTextEmbeddingClient(SpringUtil.getBean(TextEmbedding.class), connectionProperties); return new TongYiAutoConfiguration().tongYiTextEmbeddingClient(SpringUtil.getBean(TextEmbedding.class), connectionProperties);
} }
private static FastGPTChatModel buildFastGPTModel(String openAiToken, String url) {
url = StrUtil.blankToDefault(url, ApiUtils.DEFAULT_BASE_URL);
FastGPTApi openAiApi = new FastGPTApi(url, openAiToken);
return new FastGPTChatModel(openAiApi);
}
} }
package cn.iocoder.yudao.framework.ai.core.model.fastgpt;
import com.fasterxml.jackson.annotation.JsonInclude;
import com.fasterxml.jackson.annotation.JsonProperty;
import com.fasterxml.jackson.annotation.JsonInclude.Include;
import java.util.List;
import java.util.Map;
import java.util.concurrent.atomic.AtomicBoolean;
import java.util.function.Predicate;
import org.springframework.ai.model.ModelDescription;
import org.springframework.ai.model.ModelOptionsUtils;
import org.springframework.ai.openai.api.ApiUtils;
import org.springframework.ai.retry.RetryUtils;
import org.springframework.boot.context.properties.bind.ConstructorBinding;
import org.springframework.core.ParameterizedTypeReference;
import org.springframework.http.ResponseEntity;
import org.springframework.util.Assert;
import org.springframework.util.CollectionUtils;
import org.springframework.web.client.ResponseErrorHandler;
import org.springframework.web.client.RestClient;
import org.springframework.web.reactive.function.client.WebClient;
import reactor.core.publisher.Flux;
import reactor.core.publisher.Mono;
public class FastGPTApi {
public static final String DEFAULT_CHAT_MODEL;
public static final String DEFAULT_EMBEDDING_MODEL;
private static final Predicate<String> SSE_DONE_PREDICATE;
private final RestClient restClient;
private final WebClient webClient;
private FastGPTStreamFunctionCallingHelper chunkMerger;
public FastGPTApi(String openAiToken) {
this("https://api.openai.com", openAiToken);
}
public FastGPTApi(String baseUrl, String openAiToken) {
this(baseUrl, openAiToken, RestClient.builder(), WebClient.builder());
}
public FastGPTApi(String baseUrl, String openAiToken, RestClient.Builder restClientBuilder, WebClient.Builder webClientBuilder) {
this(baseUrl, openAiToken, restClientBuilder, webClientBuilder, RetryUtils.DEFAULT_RESPONSE_ERROR_HANDLER);
}
public FastGPTApi(String baseUrl, String openAiToken, RestClient.Builder restClientBuilder, WebClient.Builder webClientBuilder, ResponseErrorHandler responseErrorHandler) {
this.chunkMerger = new FastGPTStreamFunctionCallingHelper();
this.restClient = restClientBuilder.baseUrl(baseUrl).defaultHeaders(ApiUtils.getJsonContentHeaders(openAiToken)).defaultStatusHandler(responseErrorHandler).build();
this.webClient = webClientBuilder.baseUrl(baseUrl).defaultHeaders(ApiUtils.getJsonContentHeaders(openAiToken)).build();
}
public static String getTextContent(List<ChatCompletionMessage.MediaContent> content) {
return (String)content.stream().filter((c) -> {
return "text".equals(c.type());
}).map(ChatCompletionMessage.MediaContent::text).reduce("", (a, b) -> {
return a + b;
});
}
public ResponseEntity<ChatCompletion> chatCompletionEntity(ChatCompletionRequest chatRequest) {
Assert.notNull(chatRequest, "The request body can not be null.");
Assert.isTrue(!chatRequest.stream(), "Request must set the steam property to false.");
return ((RestClient.RequestBodySpec)this.restClient.post().uri("/v1/chat/completions", new Object[0])).body(chatRequest).retrieve().toEntity(ChatCompletion.class);
}
public Flux<ChatCompletionChunk> chatCompletionStream(ChatCompletionRequest chatRequest) {
Assert.notNull(chatRequest, "The request body can not be null.");
Assert.isTrue(chatRequest.stream(), "Request must set the steam property to true.");
AtomicBoolean isInsideTool = new AtomicBoolean(false);
return ((WebClient.RequestBodySpec)this.webClient.post().uri("/v1/chat/completions", new Object[0])).body(Mono.just(chatRequest), ChatCompletionRequest.class).retrieve().bodyToFlux(String.class).takeUntil(SSE_DONE_PREDICATE).filter(SSE_DONE_PREDICATE.negate()).map((content) -> {
return (ChatCompletionChunk)ModelOptionsUtils.jsonToObject(content, ChatCompletionChunk.class);
}).map((chunk) -> {
if (this.chunkMerger.isStreamingToolFunctionCall(chunk)) {
isInsideTool.set(true);
}
return chunk;
}).windowUntil((chunk) -> {
if (isInsideTool.get() && this.chunkMerger.isStreamingToolFunctionCallFinish(chunk)) {
isInsideTool.set(false);
return true;
} else {
return !isInsideTool.get();
}
}).concatMapIterable((window) -> {
Mono<ChatCompletionChunk> monoChunk = window.reduce(new ChatCompletionChunk((String)null, (List)null, (Long)null, (String)null, (String)null, (String)null), (previous, current) -> {
return this.chunkMerger.merge(previous, current);
});
return List.of(monoChunk);
}).flatMap((mono) -> {
return mono;
});
}
public <T> ResponseEntity<EmbeddingList<Embedding>> embeddings(EmbeddingRequest<T> embeddingRequest) {
Assert.notNull(embeddingRequest, "The request body can not be null.");
Assert.notNull(embeddingRequest.input(), "The input can not be null.");
Assert.isTrue(embeddingRequest.input() instanceof String || embeddingRequest.input() instanceof List, "The input must be either a String, or a List of Strings or List of List of integers.");
Object var3 = embeddingRequest.input();
if (var3 instanceof List list) {
Assert.isTrue(!CollectionUtils.isEmpty(list), "The input list can not be empty.");
Assert.isTrue(list.size() <= 2048, "The list must be 2048 dimensions or less");
Assert.isTrue(list.get(0) instanceof String || list.get(0) instanceof Integer || list.get(0) instanceof List, "The input must be either a String, or a List of Strings or list of list of integers.");
}
return ((RestClient.RequestBodySpec)this.restClient.post().uri("/v1/embeddings", new Object[0])).body(embeddingRequest).retrieve().toEntity(new ParameterizedTypeReference<EmbeddingList<Embedding>>() {
});
}
static {
DEFAULT_CHAT_MODEL = FastGPTApi.ChatModel.GPT_3_5_TURBO.getValue();
DEFAULT_EMBEDDING_MODEL = FastGPTApi.EmbeddingModel.TEXT_EMBEDDING_ADA_002.getValue();
SSE_DONE_PREDICATE = "[DONE]"::equals;
}
@JsonInclude(Include.NON_NULL)
public static record ChatCompletionRequest(List<ChatCompletionMessage> messages, String model,String chatId, Float frequencyPenalty, Map<String, Integer> logitBias, Boolean logprobs, Integer topLogprobs, Integer maxTokens, Integer n, Float presencePenalty, ResponseFormat responseFormat, Integer seed, List<String> stop, Boolean stream, Float temperature, Float topP, List<FunctionTool> tools, Object toolChoice, String user) {
public ChatCompletionRequest(List<ChatCompletionMessage> messages, String model, Float temperature) {
this(messages, model, (String)null, (Float)null, (Map)null, (Boolean)null, (Integer)null, (Integer)null, (Integer)null, (Float)null, (ResponseFormat)null, (Integer)null, (List)null, false, temperature, (Float)null, (List)null, (Object)null, (String)null);
}
public ChatCompletionRequest(List<ChatCompletionMessage> messages, String model, Float temperature, boolean stream) {
this(messages, model, (String)null, (Float)null, (Map)null, (Boolean)null, (Integer)null, (Integer)null, (Integer)null, (Float)null, (ResponseFormat)null, (Integer)null, (List)null, stream, temperature, (Float)null, (List)null, (Object)null, (String)null);
}
public ChatCompletionRequest(List<ChatCompletionMessage> messages, String model, List<FunctionTool> tools, Object toolChoice) {
this(messages, model, (String)null, (Float)null, (Map)null, (Boolean)null, (Integer)null, (Integer)null, (Integer)null, (Float)null, (ResponseFormat)null, (Integer)null, (List)null, false, 0.8F, (Float)null, tools, toolChoice, (String)null);
}
public ChatCompletionRequest(List<ChatCompletionMessage> messages, Boolean stream) {
this(messages, (String)null, (String)null, (Float)null, (Map)null, (Boolean)null, (Integer)null, (Integer)null, (Integer)null, (Float)null, (ResponseFormat)null, (Integer)null, (List)null, stream, (Float)null, (Float)null, (List)null, (Object)null, (String)null);
}
public ChatCompletionRequest(List<ChatCompletionMessage> messages, String model, String chatId, Boolean stream) {
this(messages, model, chatId, (Float)null, (Map)null, (Boolean)null, (Integer)null, (Integer)null, (Integer)null, (Float)null, (ResponseFormat)null, (Integer)null, (List)null, stream, (Float)null, (Float)null, (List)null, (Object)null, (String)null);
}
public ChatCompletionRequest(@JsonProperty("messages") List<ChatCompletionMessage> messages, @JsonProperty("model") String model, @JsonProperty("chatId") String chatId, @JsonProperty("frequency_penalty") Float frequencyPenalty, @JsonProperty("logit_bias") Map<String, Integer> logitBias, @JsonProperty("logprobs") Boolean logprobs, @JsonProperty("top_logprobs") Integer topLogprobs, @JsonProperty("max_tokens") Integer maxTokens, @JsonProperty("n") Integer n, @JsonProperty("presence_penalty") Float presencePenalty, @JsonProperty("response_format") ResponseFormat responseFormat, @JsonProperty("seed") Integer seed, @JsonProperty("stop") List<String> stop, @JsonProperty("stream") Boolean stream, @JsonProperty("temperature") Float temperature, @JsonProperty("top_p") Float topP, @JsonProperty("tools") List<FunctionTool> tools, @JsonProperty("tool_choice") Object toolChoice, @JsonProperty("user") String user) {
this.messages = messages;
this.model = model;
this.chatId = chatId;
this.frequencyPenalty = frequencyPenalty;
this.logitBias = logitBias;
this.logprobs = logprobs;
this.topLogprobs = topLogprobs;
this.maxTokens = maxTokens;
this.n = n;
this.presencePenalty = presencePenalty;
this.responseFormat = responseFormat;
this.seed = seed;
this.stop = stop;
this.stream = stream;
this.temperature = temperature;
this.topP = topP;
this.tools = tools;
this.toolChoice = toolChoice;
this.user = user;
}
@JsonProperty("messages")
public List<ChatCompletionMessage> messages() {
return this.messages;
}
@JsonProperty("model")
public String model() {
return this.model;
}
@JsonProperty("chatId")
public String chatId() {
return this.chatId;
}
@JsonProperty("frequency_penalty")
public Float frequencyPenalty() {
return this.frequencyPenalty;
}
@JsonProperty("logit_bias")
public Map<String, Integer> logitBias() {
return this.logitBias;
}
@JsonProperty("logprobs")
public Boolean logprobs() {
return this.logprobs;
}
@JsonProperty("top_logprobs")
public Integer topLogprobs() {
return this.topLogprobs;
}
@JsonProperty("max_tokens")
public Integer maxTokens() {
return this.maxTokens;
}
@JsonProperty("n")
public Integer n() {
return this.n;
}
@JsonProperty("presence_penalty")
public Float presencePenalty() {
return this.presencePenalty;
}
@JsonProperty("response_format")
public ResponseFormat responseFormat() {
return this.responseFormat;
}
@JsonProperty("seed")
public Integer seed() {
return this.seed;
}
@JsonProperty("stop")
public List<String> stop() {
return this.stop;
}
@JsonProperty("stream")
public Boolean stream() {
return this.stream;
}
@JsonProperty("temperature")
public Float temperature() {
return this.temperature;
}
@JsonProperty("top_p")
public Float topP() {
return this.topP;
}
@JsonProperty("tools")
public List<FunctionTool> tools() {
return this.tools;
}
@JsonProperty("tool_choice")
public Object toolChoice() {
return this.toolChoice;
}
@JsonProperty("user")
public String user() {
return this.user;
}
@JsonInclude(Include.NON_NULL)
public static record ResponseFormat(String type) {
public ResponseFormat(@JsonProperty("type") String type) {
this.type = type;
}
@JsonProperty("type")
public String type() {
return this.type;
}
}
public static class ToolChoiceBuilder {
public static final String AUTO = "auto";
public static final String NONE = "none";
public ToolChoiceBuilder() {
}
public static Object FUNCTION(String functionName) {
return Map.of("type", "function", "function", Map.of("name", functionName));
}
}
}
@JsonInclude(Include.NON_NULL)
public static record ChatCompletion(String id, List<Choice> choices, Long created, String model, String systemFingerprint, String object, Usage usage) {
public ChatCompletion(@JsonProperty("id") String id, @JsonProperty("choices") List<Choice> choices, @JsonProperty("created") Long created, @JsonProperty("model") String model, @JsonProperty("system_fingerprint") String systemFingerprint, @JsonProperty("object") String object, @JsonProperty("usage") Usage usage) {
this.id = id;
this.choices = choices;
this.created = created;
this.model = model;
this.systemFingerprint = systemFingerprint;
this.object = object;
this.usage = usage;
}
@JsonProperty("id")
public String id() {
return this.id;
}
@JsonProperty("choices")
public List<Choice> choices() {
return this.choices;
}
@JsonProperty("created")
public Long created() {
return this.created;
}
@JsonProperty("model")
public String model() {
return this.model;
}
@JsonProperty("system_fingerprint")
public String systemFingerprint() {
return this.systemFingerprint;
}
@JsonProperty("object")
public String object() {
return this.object;
}
@JsonProperty("usage")
public Usage usage() {
return this.usage;
}
@JsonInclude(Include.NON_NULL)
public static record Choice(ChatCompletionFinishReason finishReason, Integer index, ChatCompletionMessage message, LogProbs logprobs) {
public Choice(@JsonProperty("finish_reason") ChatCompletionFinishReason finishReason, @JsonProperty("index") Integer index, @JsonProperty("message") ChatCompletionMessage message, @JsonProperty("logprobs") LogProbs logprobs) {
this.finishReason = finishReason;
this.index = index;
this.message = message;
this.logprobs = logprobs;
}
@JsonProperty("finish_reason")
public ChatCompletionFinishReason finishReason() {
return this.finishReason;
}
@JsonProperty("index")
public Integer index() {
return this.index;
}
@JsonProperty("message")
public ChatCompletionMessage message() {
return this.message;
}
@JsonProperty("logprobs")
public LogProbs logprobs() {
return this.logprobs;
}
}
}
@JsonInclude(Include.NON_NULL)
public static record EmbeddingRequest<T>(T input, String model, String encodingFormat, Integer dimensions, String user) {
public EmbeddingRequest(T input, String model) {
this(input, model, "float", (Integer)null, (String)null);
}
public EmbeddingRequest(T input) {
this(input, FastGPTApi.DEFAULT_EMBEDDING_MODEL);
}
public EmbeddingRequest(@JsonProperty("input") T input, @JsonProperty("model") String model, @JsonProperty("encoding_format") String encodingFormat, @JsonProperty("dimensions") Integer dimensions, @JsonProperty("user") String user) {
this.input = input;
this.model = model;
this.encodingFormat = encodingFormat;
this.dimensions = dimensions;
this.user = user;
}
@JsonProperty("input")
public T input() {
return this.input;
}
@JsonProperty("model")
public String model() {
return this.model;
}
@JsonProperty("encoding_format")
public String encodingFormat() {
return this.encodingFormat;
}
@JsonProperty("dimensions")
public Integer dimensions() {
return this.dimensions;
}
@JsonProperty("user")
public String user() {
return this.user;
}
}
@JsonInclude(Include.NON_NULL)
public static record ChatCompletionChunk(String id, List<ChunkChoice> choices, Long created, String model, String systemFingerprint, String object) {
public ChatCompletionChunk(@JsonProperty("id") String id, @JsonProperty("choices") List<ChunkChoice> choices, @JsonProperty("created") Long created, @JsonProperty("model") String model, @JsonProperty("system_fingerprint") String systemFingerprint, @JsonProperty("object") String object) {
this.id = id;
this.choices = choices;
this.created = created;
this.model = model;
this.systemFingerprint = systemFingerprint;
this.object = object;
}
@JsonProperty("id")
public String id() {
return this.id;
}
@JsonProperty("choices")
public List<ChunkChoice> choices() {
return this.choices;
}
@JsonProperty("created")
public Long created() {
return this.created;
}
@JsonProperty("model")
public String model() {
return this.model;
}
@JsonProperty("system_fingerprint")
public String systemFingerprint() {
return this.systemFingerprint;
}
@JsonProperty("object")
public String object() {
return this.object;
}
@JsonInclude(Include.NON_NULL)
public static record ChunkChoice(ChatCompletionFinishReason finishReason, Integer index, ChatCompletionMessage delta, LogProbs logprobs) {
public ChunkChoice(@JsonProperty("finish_reason") ChatCompletionFinishReason finishReason, @JsonProperty("index") Integer index, @JsonProperty("delta") ChatCompletionMessage delta, @JsonProperty("logprobs") LogProbs logprobs) {
this.finishReason = finishReason;
this.index = index;
this.delta = delta;
this.logprobs = logprobs;
}
@JsonProperty("finish_reason")
public ChatCompletionFinishReason finishReason() {
return this.finishReason;
}
@JsonProperty("index")
public Integer index() {
return this.index;
}
@JsonProperty("delta")
public ChatCompletionMessage delta() {
return this.delta;
}
@JsonProperty("logprobs")
public LogProbs logprobs() {
return this.logprobs;
}
}
}
@JsonInclude(Include.NON_NULL)
public static record ChatCompletionMessage(Object rawContent, Role role, String name, String toolCallId, List<ToolCall> toolCalls) {
public ChatCompletionMessage(Object content, Role role) {
this(content, role, (String)null, (String)null, (List)null);
}
public ChatCompletionMessage(@JsonProperty("content") Object rawContent, @JsonProperty("role") Role role, @JsonProperty("name") String name, @JsonProperty("tool_call_id") String toolCallId, @JsonProperty("tool_calls") List<ToolCall> toolCalls) {
this.rawContent = rawContent;
this.role = role;
this.name = name;
this.toolCallId = toolCallId;
this.toolCalls = toolCalls;
}
public String content() {
if (this.rawContent == null) {
return null;
} else {
Object var2 = this.rawContent;
if (var2 instanceof String) {
String text = (String)var2;
return text;
} else {
throw new IllegalStateException("The content is not a string!");
}
}
}
@JsonProperty("content")
public Object rawContent() {
return this.rawContent;
}
@JsonProperty("role")
public Role role() {
return this.role;
}
@JsonProperty("name")
public String name() {
return this.name;
}
@JsonProperty("tool_call_id")
public String toolCallId() {
return this.toolCallId;
}
@JsonProperty("tool_calls")
public List<ToolCall> toolCalls() {
return this.toolCalls;
}
public static enum Role {
@JsonProperty("system")
SYSTEM,
@JsonProperty("user")
USER,
@JsonProperty("assistant")
ASSISTANT,
@JsonProperty("tool")
TOOL;
private Role() {
}
}
@JsonInclude(Include.NON_NULL)
public static record ChatCompletionFunction(String name, String arguments) {
public ChatCompletionFunction(@JsonProperty("name") String name, @JsonProperty("arguments") String arguments) {
this.name = name;
this.arguments = arguments;
}
@JsonProperty("name")
public String name() {
return this.name;
}
@JsonProperty("arguments")
public String arguments() {
return this.arguments;
}
}
@JsonInclude(Include.NON_NULL)
public static record ToolCall(String id, String type, ChatCompletionFunction function) {
public ToolCall(@JsonProperty("id") String id, @JsonProperty("type") String type, @JsonProperty("function") ChatCompletionFunction function) {
this.id = id;
this.type = type;
this.function = function;
}
@JsonProperty("id")
public String id() {
return this.id;
}
@JsonProperty("type")
public String type() {
return this.type;
}
@JsonProperty("function")
public ChatCompletionFunction function() {
return this.function;
}
}
@JsonInclude(Include.NON_NULL)
public static record MediaContent(String type, String text, ImageUrl imageUrl) {
public MediaContent(String text) {
this("text", text, (ImageUrl)null);
}
public MediaContent(ImageUrl imageUrl) {
this("image_url", (String)null, imageUrl);
}
public MediaContent(@JsonProperty("type") String type, @JsonProperty("text") String text, @JsonProperty("image_url") ImageUrl imageUrl) {
this.type = type;
this.text = text;
this.imageUrl = imageUrl;
}
@JsonProperty("type")
public String type() {
return this.type;
}
@JsonProperty("text")
public String text() {
return this.text;
}
@JsonProperty("image_url")
public ImageUrl imageUrl() {
return this.imageUrl;
}
@JsonInclude(Include.NON_NULL)
public static record ImageUrl(String url, String detail) {
public ImageUrl(String url) {
this(url, (String)null);
}
public ImageUrl(@JsonProperty("url") String url, @JsonProperty("detail") String detail) {
this.url = url;
this.detail = detail;
}
@JsonProperty("url")
public String url() {
return this.url;
}
@JsonProperty("detail")
public String detail() {
return this.detail;
}
}
}
}
public static enum ChatModel implements ModelDescription {
GPT_4_O("gpt-4o"),
GPT_4_TURBO("gpt-4-turbo"),
GPT_4_TURBO_2204_04_09("gpt-4-turbo-2024-04-09"),
GPT_4_0125_PREVIEW("gpt-4-0125-preview"),
GPT_4_TURBO_PREVIEW("gpt-4-turbo-preview"),
GPT_4_VISION_PREVIEW("gpt-4-vision-preview"),
GPT_4("gpt-4"),
GPT_4_32K("gpt-4-32k"),
GPT_3_5_TURBO("gpt-3.5-turbo"),
GPT_3_5_TURBO_0125("gpt-3.5-turbo-0125"),
GPT_3_5_TURBO_1106("gpt-3.5-turbo-1106");
public final String value;
private ChatModel(String value) {
this.value = value;
}
public String getValue() {
return this.value;
}
public String getModelName() {
return this.value;
}
}
public static enum EmbeddingModel {
TEXT_EMBEDDING_3_LARGE("text-embedding-3-large"),
TEXT_EMBEDDING_3_SMALL("text-embedding-3-small"),
TEXT_EMBEDDING_ADA_002("text-embedding-ada-002");
public final String value;
private EmbeddingModel(String value) {
this.value = value;
}
public String getValue() {
return this.value;
}
}
@JsonInclude(Include.NON_NULL)
public static record EmbeddingList<T>(String object, List<T> data, String model, Usage usage) {
public EmbeddingList(@JsonProperty("object") String object, @JsonProperty("data") List<T> data, @JsonProperty("model") String model, @JsonProperty("usage") Usage usage) {
this.object = object;
this.data = data;
this.model = model;
this.usage = usage;
}
@JsonProperty("object")
public String object() {
return this.object;
}
@JsonProperty("data")
public List<T> data() {
return this.data;
}
@JsonProperty("model")
public String model() {
return this.model;
}
@JsonProperty("usage")
public Usage usage() {
return this.usage;
}
}
@JsonInclude(Include.NON_NULL)
public static record Embedding(Integer index, List<Double> embedding, String object) {
public Embedding(Integer index, List<Double> embedding) {
this(index, embedding, "embedding");
}
public Embedding(@JsonProperty("index") Integer index, @JsonProperty("embedding") List<Double> embedding, @JsonProperty("object") String object) {
this.index = index;
this.embedding = embedding;
this.object = object;
}
@JsonProperty("index")
public Integer index() {
return this.index;
}
@JsonProperty("embedding")
public List<Double> embedding() {
return this.embedding;
}
@JsonProperty("object")
public String object() {
return this.object;
}
}
@JsonInclude(Include.NON_NULL)
public static record Usage(Integer completionTokens, Integer promptTokens, Integer totalTokens) {
public Usage(@JsonProperty("completion_tokens") Integer completionTokens, @JsonProperty("prompt_tokens") Integer promptTokens, @JsonProperty("total_tokens") Integer totalTokens) {
this.completionTokens = completionTokens;
this.promptTokens = promptTokens;
this.totalTokens = totalTokens;
}
@JsonProperty("completion_tokens")
public Integer completionTokens() {
return this.completionTokens;
}
@JsonProperty("prompt_tokens")
public Integer promptTokens() {
return this.promptTokens;
}
@JsonProperty("total_tokens")
public Integer totalTokens() {
return this.totalTokens;
}
}
@JsonInclude(Include.NON_NULL)
public static record LogProbs(List<Content> content) {
public LogProbs(@JsonProperty("content") List<Content> content) {
this.content = content;
}
@JsonProperty("content")
public List<Content> content() {
return this.content;
}
@JsonInclude(Include.NON_NULL)
public static record Content(String token, Float logprob, List<Integer> probBytes, List<TopLogProbs> topLogprobs) {
public Content(@JsonProperty("token") String token, @JsonProperty("logprob") Float logprob, @JsonProperty("bytes") List<Integer> probBytes, @JsonProperty("top_logprobs") List<TopLogProbs> topLogprobs) {
this.token = token;
this.logprob = logprob;
this.probBytes = probBytes;
this.topLogprobs = topLogprobs;
}
@JsonProperty("token")
public String token() {
return this.token;
}
@JsonProperty("logprob")
public Float logprob() {
return this.logprob;
}
@JsonProperty("bytes")
public List<Integer> probBytes() {
return this.probBytes;
}
@JsonProperty("top_logprobs")
public List<TopLogProbs> topLogprobs() {
return this.topLogprobs;
}
@JsonInclude(Include.NON_NULL)
public static record TopLogProbs(String token, Float logprob, List<Integer> probBytes) {
public TopLogProbs(@JsonProperty("token") String token, @JsonProperty("logprob") Float logprob, @JsonProperty("bytes") List<Integer> probBytes) {
this.token = token;
this.logprob = logprob;
this.probBytes = probBytes;
}
@JsonProperty("token")
public String token() {
return this.token;
}
@JsonProperty("logprob")
public Float logprob() {
return this.logprob;
}
@JsonProperty("bytes")
public List<Integer> probBytes() {
return this.probBytes;
}
}
}
}
public static enum ChatCompletionFinishReason {
@JsonProperty("stop")
STOP,
@JsonProperty("length")
LENGTH,
@JsonProperty("content_filter")
CONTENT_FILTER,
@JsonProperty("tool_calls")
TOOL_CALLS,
@JsonProperty("function_call")
FUNCTION_CALL,
@JsonProperty("tool_call")
TOOL_CALL;
private ChatCompletionFinishReason() {
}
}
@JsonInclude(Include.NON_NULL)
public static record FunctionTool(Type type, Function function) {
@ConstructorBinding
public FunctionTool(Function function) {
this(FastGPTApi.FunctionTool.Type.FUNCTION, function);
}
public FunctionTool(@JsonProperty("type") Type type, @JsonProperty("function") Function function) {
this.type = type;
this.function = function;
}
@JsonProperty("type")
public Type type() {
return this.type;
}
@JsonProperty("function")
public Function function() {
return this.function;
}
public static enum Type {
@JsonProperty("function")
FUNCTION;
private Type() {
}
}
public static record Function(String description, String name, Map<String, Object> parameters) {
@ConstructorBinding
public Function(String description, String name, String jsonSchema) {
this(description, name, ModelOptionsUtils.jsonToMap(jsonSchema));
}
public Function(@JsonProperty("description") String description, @JsonProperty("name") String name, @JsonProperty("parameters") Map<String, Object> parameters) {
this.description = description;
this.name = name;
this.parameters = parameters;
}
@JsonProperty("description")
public String description() {
return this.description;
}
@JsonProperty("name")
public String name() {
return this.name;
}
@JsonProperty("parameters")
public Map<String, Object> parameters() {
return this.parameters;
}
}
}
}
package cn.iocoder.yudao.framework.ai.core.model.fastgpt;
import java.util.ArrayList;
import java.util.Base64;
import java.util.HashMap;
import java.util.HashSet;
import java.util.Iterator;
import java.util.List;
import java.util.Map;
import java.util.Optional;
import java.util.Set;
import java.util.concurrent.ConcurrentHashMap;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import org.springframework.ai.chat.metadata.ChatGenerationMetadata;
import org.springframework.ai.chat.metadata.RateLimit;
import org.springframework.ai.chat.model.ChatModel;
import org.springframework.ai.chat.model.ChatResponse;
import org.springframework.ai.chat.model.Generation;
import org.springframework.ai.chat.model.StreamingChatModel;
import org.springframework.ai.chat.prompt.ChatOptions;
import org.springframework.ai.chat.prompt.Prompt;
import org.springframework.ai.model.ModelOptions;
import org.springframework.ai.model.ModelOptionsUtils;
import org.springframework.ai.model.function.AbstractFunctionCallSupport;
import org.springframework.ai.model.function.FunctionCallback;
import org.springframework.ai.model.function.FunctionCallbackContext;
import cn.iocoder.yudao.framework.ai.core.model.fastgpt.FastGPTApi;
import cn.iocoder.yudao.framework.ai.core.model.fastgpt.FastGPTApi.ChatCompletionFinishReason;
import cn.iocoder.yudao.framework.ai.core.model.fastgpt.FastGPTApi.ChatCompletionMessage.Role;
import org.springframework.ai.openai.api.OpenAiApi;
import org.springframework.ai.openai.metadata.OpenAiChatResponseMetadata;
import org.springframework.ai.openai.metadata.support.OpenAiResponseHeaderExtractor;
import org.springframework.ai.retry.RetryUtils;
import org.springframework.http.HttpEntity;
import org.springframework.http.ResponseEntity;
import org.springframework.retry.support.RetryTemplate;
import org.springframework.util.Assert;
import org.springframework.util.CollectionUtils;
import org.springframework.util.MimeType;
import reactor.core.publisher.Flux;
public class FastGPTChatModel extends AbstractFunctionCallSupport<FastGPTApi.ChatCompletionMessage, FastGPTApi.ChatCompletionRequest, ResponseEntity<FastGPTApi.ChatCompletion>> implements ChatModel, StreamingChatModel {
private static final Logger logger = LoggerFactory.getLogger(FastGPTChatModel.class);
private FastGPTOptions defaultOptions;
private final RetryTemplate retryTemplate;
private final FastGPTApi fastGPTApi;
public FastGPTChatModel(FastGPTApi fastGPTApi) {
this(fastGPTApi, FastGPTOptions.builder().withModel(FastGPTApi.DEFAULT_CHAT_MODEL).withTemperature(0.7F).build());
}
public FastGPTChatModel(FastGPTApi fastGPTApi, FastGPTOptions options) {
this(fastGPTApi, options, (FunctionCallbackContext)null, RetryUtils.DEFAULT_RETRY_TEMPLATE);
}
public FastGPTChatModel(FastGPTApi fastGPTApi, FastGPTOptions options, FunctionCallbackContext functionCallbackContext, RetryTemplate retryTemplate) {
super(functionCallbackContext);
Assert.notNull(fastGPTApi, "FastGPTApi must not be null");
Assert.notNull(options, "Options must not be null");
Assert.notNull(retryTemplate, "RetryTemplate must not be null");
this.fastGPTApi = fastGPTApi;
this.defaultOptions = options;
this.retryTemplate = retryTemplate;
}
public ChatResponse call(Prompt prompt) {
FastGPTApi.ChatCompletionRequest request = this.createRequest(prompt, false);
return (ChatResponse)this.retryTemplate.execute((ctx) -> {
ResponseEntity<FastGPTApi.ChatCompletion> completionEntity = (ResponseEntity)this.callWithFunctionSupport(request);
FastGPTApi.ChatCompletion chatCompletion = (FastGPTApi.ChatCompletion)completionEntity.getBody();
if (chatCompletion == null) {
logger.warn("No chat completion returned for prompt: {}", prompt);
return new ChatResponse(List.of());
} else {
RateLimit rateLimits = OpenAiResponseHeaderExtractor.extractAiResponseHeaders(completionEntity);
List<FastGPTApi.ChatCompletion.Choice> choices = chatCompletion.choices();
if (choices == null) {
logger.warn("No choices returned for prompt: {}", prompt);
return new ChatResponse(List.of());
} else {
List<Generation> generations = choices.stream().map((choice) -> {
return (new Generation(choice.message().content(), this.toMap(chatCompletion.id(), choice))).withGenerationMetadata(ChatGenerationMetadata.from(choice.finishReason().name(), (Object)null));
}).toList();
return new ChatResponse(generations, FastGPTResponseMetadata.from((FastGPTApi.ChatCompletion)completionEntity.getBody()).withRateLimit(rateLimits));
}
}
});
}
private Map<String, Object> toMap(String id, FastGPTApi.ChatCompletion.Choice choice) {
Map<String, Object> map = new HashMap();
FastGPTApi.ChatCompletionMessage message = choice.message();
if (message.role() != null) {
map.put("role", message.role().name());
}
if (choice.finishReason() != null) {
map.put("finishReason", choice.finishReason().name());
}
map.put("id", id);
return map;
}
public Flux<ChatResponse> stream(Prompt prompt) {
FastGPTApi.ChatCompletionRequest request = this.createRequest(prompt, true);
return (Flux)this.retryTemplate.execute((ctx) -> {
Flux<FastGPTApi.ChatCompletionChunk> completionChunks = this.fastGPTApi.chatCompletionStream(request);
ConcurrentHashMap<String, String> roleMap = new ConcurrentHashMap();
return completionChunks.map((chunk) -> {
return this.chunkToChatCompletion(chunk);
}).switchMap((cc) -> {
return this.handleFunctionCallOrReturnStream(request, Flux.just(ResponseEntity.of(Optional.of(cc))));
}).map(HttpEntity::getBody).map((chatCompletion) -> {
try {
String id = chatCompletion.id();
List<Generation> generations = chatCompletion.choices().stream().map((choice) -> {
if (choice.message().role() != null) {
roleMap.putIfAbsent(id, choice.message().role().name());
}
String finish = choice.finishReason() != null ? choice.finishReason().name() : "";
Generation generation = new Generation(choice.message().content(), Map.of("id", id, "role", roleMap.get(id), "finishReason", finish));
if (choice.finishReason() != null) {
generation = generation.withGenerationMetadata(ChatGenerationMetadata.from(choice.finishReason().name(), (Object)null));
}
return generation;
}).toList();
return new ChatResponse(generations);
} catch (Exception var4) {
logger.error("Error processing chat completion", var4);
return new ChatResponse(List.of());
}
});
});
}
private FastGPTApi.ChatCompletion chunkToChatCompletion(FastGPTApi.ChatCompletionChunk chunk) {
List<FastGPTApi.ChatCompletion.Choice> choices = chunk.choices().stream().map((cc) -> {
return new FastGPTApi.ChatCompletion.Choice(cc.finishReason(), cc.index(), cc.delta(), cc.logprobs());
}).toList();
return new FastGPTApi.ChatCompletion(chunk.id(), choices, chunk.created(), chunk.model(), chunk.systemFingerprint(), "chat.completion", (FastGPTApi.Usage)null);
}
FastGPTApi.ChatCompletionRequest createRequest(Prompt prompt, boolean stream) {
Set<String> functionsForThisRequest = new HashSet();
List<FastGPTApi.ChatCompletionMessage> chatCompletionMessages = prompt.getInstructions().stream().map((m) -> {
List<FastGPTApi.ChatCompletionMessage.MediaContent> contents = new ArrayList(List.of(new FastGPTApi.ChatCompletionMessage.MediaContent(m.getContent())));
if (!CollectionUtils.isEmpty(m.getMedia())) {
contents.addAll(m.getMedia().stream().map((media) -> {
return new FastGPTApi.ChatCompletionMessage.MediaContent(new FastGPTApi.ChatCompletionMessage.MediaContent.ImageUrl(this.fromMediaData(media.getMimeType(), media.getData())));
}).toList());
}
return new FastGPTApi.ChatCompletionMessage(contents, Role.valueOf(m.getMessageType().name()));
}).toList();
FastGPTApi.ChatCompletionRequest request = new FastGPTApi.ChatCompletionRequest(chatCompletionMessages, stream);
if (prompt.getOptions() != null) {
ModelOptions var7 = prompt.getOptions();
if (!(var7 instanceof ChatOptions)) {
throw new IllegalArgumentException("Prompt options are not of type ChatOptions: " + prompt.getOptions().getClass().getSimpleName());
}
ChatOptions runtimeOptions = (ChatOptions)var7;
FastGPTOptions updatedRuntimeOptions = (FastGPTOptions)ModelOptionsUtils.copyToTarget(runtimeOptions, ChatOptions.class, FastGPTOptions.class);
Set<String> promptEnabledFunctions = this.handleFunctionCallbackConfigurations(updatedRuntimeOptions, true);
functionsForThisRequest.addAll(promptEnabledFunctions);
// request = (FastGPTApi.ChatCompletionRequest)ModelOptionsUtils.merge(updatedRuntimeOptions, request, FastGPTApi.ChatCompletionRequest.class);
request = new FastGPTApi.ChatCompletionRequest(chatCompletionMessages, updatedRuntimeOptions.getModel(), updatedRuntimeOptions.getChatId(), stream);
}
// if (this.defaultOptions != null) {
// Set<String> defaultEnabledFunctions = this.handleFunctionCallbackConfigurations(this.defaultOptions, false);
// functionsForThisRequest.addAll(defaultEnabledFunctions);
// request = (FastGPTApi.ChatCompletionRequest)ModelOptionsUtils.merge(request, this.defaultOptions, FastGPTApi.ChatCompletionRequest.class);
// }
//
// if (!CollectionUtils.isEmpty(functionsForThisRequest)) {
// request = (FastGPTApi.ChatCompletionRequest)ModelOptionsUtils.merge(FastGPTOptions.builder().withTools(this.getFunctionTools(functionsForThisRequest)).build(), request, FastGPTApi.ChatCompletionRequest.class);
// }
return request;
}
private String fromMediaData(MimeType mimeType, Object mediaContentData) {
if (mediaContentData instanceof byte[] bytes) {
return String.format("data:%s;base64,%s", mimeType.toString(), Base64.getEncoder().encodeToString(bytes));
} else if (mediaContentData instanceof String text) {
return text;
} else {
throw new IllegalArgumentException("Unsupported media data type: " + mediaContentData.getClass().getSimpleName());
}
}
private List<OpenAiApi.FunctionTool> getFunctionTools(Set<String> functionNames) {
return this.resolveFunctionCallbacks(functionNames).stream().map((functionCallback) -> {
OpenAiApi.FunctionTool.Function function = new OpenAiApi.FunctionTool.Function(functionCallback.getDescription(), functionCallback.getName(), functionCallback.getInputTypeSchema());
return new OpenAiApi.FunctionTool(function);
}).toList();
}
protected FastGPTApi.ChatCompletionRequest doCreateToolResponseRequest(FastGPTApi.ChatCompletionRequest previousRequest, FastGPTApi.ChatCompletionMessage responseMessage, List<FastGPTApi.ChatCompletionMessage> conversationHistory) {
Iterator var4 = responseMessage.toolCalls().iterator();
while(var4.hasNext()) {
FastGPTApi.ChatCompletionMessage.ToolCall toolCall = (FastGPTApi.ChatCompletionMessage.ToolCall)var4.next();
String functionName = toolCall.function().name();
String functionArguments = toolCall.function().arguments();
if (!this.functionCallbackRegister.containsKey(functionName)) {
throw new IllegalStateException("No function callback found for function name: " + functionName);
}
String functionResponse = ((FunctionCallback)this.functionCallbackRegister.get(functionName)).call(functionArguments);
conversationHistory.add(new FastGPTApi.ChatCompletionMessage(functionResponse, Role.TOOL, functionName, toolCall.id(), (List)null));
}
FastGPTApi.ChatCompletionRequest newRequest = new FastGPTApi.ChatCompletionRequest(conversationHistory, previousRequest.stream());
newRequest = (FastGPTApi.ChatCompletionRequest)ModelOptionsUtils.merge(newRequest, previousRequest, FastGPTApi.ChatCompletionRequest.class);
return newRequest;
}
protected List<FastGPTApi.ChatCompletionMessage> doGetUserMessages(FastGPTApi.ChatCompletionRequest request) {
return request.messages();
}
protected FastGPTApi.ChatCompletionMessage doGetToolResponseMessage(ResponseEntity<FastGPTApi.ChatCompletion> chatCompletion) {
return ((FastGPTApi.ChatCompletion.Choice)((FastGPTApi.ChatCompletion)chatCompletion.getBody()).choices().iterator().next()).message();
}
protected ResponseEntity<FastGPTApi.ChatCompletion> doChatCompletion(FastGPTApi.ChatCompletionRequest request) {
return this.fastGPTApi.chatCompletionEntity(request);
}
protected Flux<ResponseEntity<FastGPTApi.ChatCompletion>> doChatCompletionStream(FastGPTApi.ChatCompletionRequest request) {
return this.fastGPTApi.chatCompletionStream(request).map(this::chunkToChatCompletion).map(Optional::ofNullable).map(ResponseEntity::of);
}
protected boolean isToolFunctionCall(ResponseEntity<FastGPTApi.ChatCompletion> chatCompletion) {
FastGPTApi.ChatCompletion body = (FastGPTApi.ChatCompletion)chatCompletion.getBody();
if (body == null) {
return false;
} else {
List<FastGPTApi.ChatCompletion.Choice> choices = body.choices();
if (CollectionUtils.isEmpty(choices)) {
return false;
} else {
FastGPTApi.ChatCompletion.Choice choice = (FastGPTApi.ChatCompletion.Choice)choices.get(0);
return !CollectionUtils.isEmpty(choice.message().toolCalls()) && choice.finishReason() == ChatCompletionFinishReason.TOOL_CALLS;
}
}
}
public ChatOptions getDefaultOptions() {
return FastGPTOptions.fromOptions(this.defaultOptions);
}
}
package cn.iocoder.yudao.framework.ai.core.model.fastgpt;
import com.fasterxml.jackson.annotation.JsonIgnore;
import com.fasterxml.jackson.annotation.JsonInclude;
import com.fasterxml.jackson.annotation.JsonProperty;
import com.fasterxml.jackson.annotation.JsonInclude.Include;
import java.util.ArrayList;
import java.util.HashSet;
import java.util.List;
import java.util.Map;
import java.util.Set;
import org.springframework.ai.chat.prompt.ChatOptions;
import org.springframework.ai.model.function.FunctionCallback;
import org.springframework.ai.model.function.FunctionCallingOptions;
import org.springframework.ai.openai.api.OpenAiApi;
import org.springframework.boot.context.properties.NestedConfigurationProperty;
import org.springframework.util.Assert;
@JsonInclude(Include.NON_NULL)
public class FastGPTOptions implements FunctionCallingOptions, ChatOptions {
@JsonProperty("model")
private String model;
@JsonProperty("chatId")
private String chatId;
@JsonProperty("frequency_penalty")
private Float frequencyPenalty;
@JsonProperty("logit_bias")
private Map<String, Integer> logitBias;
@JsonProperty("logprobs")
private Boolean logprobs;
@JsonProperty("top_logprobs")
private Integer topLogprobs;
@JsonProperty("max_tokens")
private Integer maxTokens;
@JsonProperty("n")
private Integer n;
@JsonProperty("presence_penalty")
private Float presencePenalty;
@JsonProperty("response_format")
private OpenAiApi.ChatCompletionRequest.ResponseFormat responseFormat;
@JsonProperty("seed")
private Integer seed;
@NestedConfigurationProperty
@JsonProperty("stop")
private List<String> stop;
@JsonProperty("temperature")
private Float temperature;
@JsonProperty("top_p")
private Float topP;
@NestedConfigurationProperty
@JsonProperty("tools")
private List<OpenAiApi.FunctionTool> tools;
@JsonProperty("tool_choice")
private String toolChoice;
@JsonProperty("user")
private String user;
@NestedConfigurationProperty
@JsonIgnore
private List<FunctionCallback> functionCallbacks = new ArrayList();
@NestedConfigurationProperty
@JsonIgnore
private Set<String> functions = new HashSet();
public FastGPTOptions() {
}
public static Builder builder() {
return new Builder();
}
public String getModel() {
return this.model;
}
public void setModel(String model) {
this.model = model;
}
public String getChatId() {
return this.chatId;
}
public void setChatId(String chatId) {
this.chatId = chatId;
}
public Float getFrequencyPenalty() {
return this.frequencyPenalty;
}
public void setFrequencyPenalty(Float frequencyPenalty) {
this.frequencyPenalty = frequencyPenalty;
}
public Map<String, Integer> getLogitBias() {
return this.logitBias;
}
public void setLogitBias(Map<String, Integer> logitBias) {
this.logitBias = logitBias;
}
public Boolean getLogprobs() {
return this.logprobs;
}
public void setLogprobs(Boolean logprobs) {
this.logprobs = logprobs;
}
public Integer getTopLogprobs() {
return this.topLogprobs;
}
public void setTopLogprobs(Integer topLogprobs) {
this.topLogprobs = topLogprobs;
}
public Integer getMaxTokens() {
return this.maxTokens;
}
public void setMaxTokens(Integer maxTokens) {
this.maxTokens = maxTokens;
}
public Integer getN() {
return this.n;
}
public void setN(Integer n) {
this.n = n;
}
public Float getPresencePenalty() {
return this.presencePenalty;
}
public void setPresencePenalty(Float presencePenalty) {
this.presencePenalty = presencePenalty;
}
public OpenAiApi.ChatCompletionRequest.ResponseFormat getResponseFormat() {
return this.responseFormat;
}
public void setResponseFormat(OpenAiApi.ChatCompletionRequest.ResponseFormat responseFormat) {
this.responseFormat = responseFormat;
}
public Integer getSeed() {
return this.seed;
}
public void setSeed(Integer seed) {
this.seed = seed;
}
public List<String> getStop() {
return this.stop;
}
public void setStop(List<String> stop) {
this.stop = stop;
}
public Float getTemperature() {
return this.temperature;
}
public void setTemperature(Float temperature) {
this.temperature = temperature;
}
public Float getTopP() {
return this.topP;
}
public void setTopP(Float topP) {
this.topP = topP;
}
public List<OpenAiApi.FunctionTool> getTools() {
return this.tools;
}
public void setTools(List<OpenAiApi.FunctionTool> tools) {
this.tools = tools;
}
public String getToolChoice() {
return this.toolChoice;
}
public void setToolChoice(String toolChoice) {
this.toolChoice = toolChoice;
}
public String getUser() {
return this.user;
}
public void setUser(String user) {
this.user = user;
}
public List<FunctionCallback> getFunctionCallbacks() {
return this.functionCallbacks;
}
public void setFunctionCallbacks(List<FunctionCallback> functionCallbacks) {
this.functionCallbacks = functionCallbacks;
}
public Set<String> getFunctions() {
return this.functions;
}
public void setFunctions(Set<String> functionNames) {
this.functions = functionNames;
}
public int hashCode() {
boolean prime = true;
int result = 1;
result = 31 * result + (this.model == null ? 0 : this.model.hashCode());
result = 31 * result + (this.chatId == null ? 0 : this.chatId.hashCode());
result = 31 * result + (this.frequencyPenalty == null ? 0 : this.frequencyPenalty.hashCode());
result = 31 * result + (this.logitBias == null ? 0 : this.logitBias.hashCode());
result = 31 * result + (this.logprobs == null ? 0 : this.logprobs.hashCode());
result = 31 * result + (this.topLogprobs == null ? 0 : this.topLogprobs.hashCode());
result = 31 * result + (this.maxTokens == null ? 0 : this.maxTokens.hashCode());
result = 31 * result + (this.n == null ? 0 : this.n.hashCode());
result = 31 * result + (this.presencePenalty == null ? 0 : this.presencePenalty.hashCode());
result = 31 * result + (this.responseFormat == null ? 0 : this.responseFormat.hashCode());
result = 31 * result + (this.seed == null ? 0 : this.seed.hashCode());
result = 31 * result + (this.stop == null ? 0 : this.stop.hashCode());
result = 31 * result + (this.temperature == null ? 0 : this.temperature.hashCode());
result = 31 * result + (this.topP == null ? 0 : this.topP.hashCode());
result = 31 * result + (this.tools == null ? 0 : this.tools.hashCode());
result = 31 * result + (this.toolChoice == null ? 0 : this.toolChoice.hashCode());
result = 31 * result + (this.user == null ? 0 : this.user.hashCode());
return result;
}
public boolean equals(Object obj) {
if (this == obj) {
return true;
} else if (obj == null) {
return false;
} else if (this.getClass() != obj.getClass()) {
return false;
} else {
FastGPTOptions other = (FastGPTOptions)obj;
if (this.model == null) {
if (other.model != null) {
return false;
}
} else if (!this.model.equals(other.model)) {
return false;
}
if (this.chatId == null) {
if (other.chatId != null) {
return false;
}
} else if (!this.chatId.equals(other.chatId)) {
return false;
}
if (this.frequencyPenalty == null) {
if (other.frequencyPenalty != null) {
return false;
}
} else if (!this.frequencyPenalty.equals(other.frequencyPenalty)) {
return false;
}
if (this.logitBias == null) {
if (other.logitBias != null) {
return false;
}
} else if (!this.logitBias.equals(other.logitBias)) {
return false;
}
if (this.logprobs == null) {
if (other.logprobs != null) {
return false;
}
} else if (!this.logprobs.equals(other.logprobs)) {
return false;
}
if (this.topLogprobs == null) {
if (other.topLogprobs != null) {
return false;
}
} else if (!this.topLogprobs.equals(other.topLogprobs)) {
return false;
}
if (this.maxTokens == null) {
if (other.maxTokens != null) {
return false;
}
} else if (!this.maxTokens.equals(other.maxTokens)) {
return false;
}
if (this.n == null) {
if (other.n != null) {
return false;
}
} else if (!this.n.equals(other.n)) {
return false;
}
if (this.presencePenalty == null) {
if (other.presencePenalty != null) {
return false;
}
} else if (!this.presencePenalty.equals(other.presencePenalty)) {
return false;
}
if (this.responseFormat == null) {
if (other.responseFormat != null) {
return false;
}
} else if (!this.responseFormat.equals(other.responseFormat)) {
return false;
}
if (this.seed == null) {
if (other.seed != null) {
return false;
}
} else if (!this.seed.equals(other.seed)) {
return false;
}
if (this.stop == null) {
if (other.stop != null) {
return false;
}
} else if (!this.stop.equals(other.stop)) {
return false;
}
if (this.temperature == null) {
if (other.temperature != null) {
return false;
}
} else if (!this.temperature.equals(other.temperature)) {
return false;
}
if (this.topP == null) {
if (other.topP != null) {
return false;
}
} else if (!this.topP.equals(other.topP)) {
return false;
}
if (this.tools == null) {
if (other.tools != null) {
return false;
}
} else if (!this.tools.equals(other.tools)) {
return false;
}
if (this.toolChoice == null) {
if (other.toolChoice != null) {
return false;
}
} else if (!this.toolChoice.equals(other.toolChoice)) {
return false;
}
if (this.user == null) {
if (other.user != null) {
return false;
}
} else if (!this.user.equals(other.user)) {
return false;
}
return true;
}
}
@JsonIgnore
public Integer getTopK() {
throw new UnsupportedOperationException("Unimplemented method 'getTopK'");
}
@JsonIgnore
public void setTopK(Integer topK) {
throw new UnsupportedOperationException("Unimplemented method 'setTopK'");
}
public static FastGPTOptions fromOptions(FastGPTOptions fromOptions) {
return builder().withModel(fromOptions.getModel()).withFrequencyPenalty(fromOptions.getFrequencyPenalty()).withLogitBias(fromOptions.getLogitBias()).withLogprobs(fromOptions.getLogprobs()).withTopLogprobs(fromOptions.getTopLogprobs()).withMaxTokens(fromOptions.getMaxTokens()).withN(fromOptions.getN()).withPresencePenalty(fromOptions.getPresencePenalty()).withResponseFormat(fromOptions.getResponseFormat()).withSeed(fromOptions.getSeed()).withStop(fromOptions.getStop()).withTemperature(fromOptions.getTemperature()).withTopP(fromOptions.getTopP()).withTools(fromOptions.getTools()).withToolChoice(fromOptions.getToolChoice()).withUser(fromOptions.getUser()).withFunctionCallbacks(fromOptions.getFunctionCallbacks()).withFunctions(fromOptions.getFunctions()).build();
}
public static class Builder {
protected FastGPTOptions options;
public Builder() {
this.options = new FastGPTOptions();
}
public Builder(FastGPTOptions options) {
this.options = options;
}
public Builder withModel(String model) {
this.options.model = model;
return this;
}
public Builder withModel(OpenAiApi.ChatModel openAiChatModel) {
this.options.model = openAiChatModel.getModelName();
return this;
}
public Builder withChatId(String chatId) {
this.options.chatId = chatId;
return this;
}
public Builder withFrequencyPenalty(Float frequencyPenalty) {
this.options.frequencyPenalty = frequencyPenalty;
return this;
}
public Builder withLogitBias(Map<String, Integer> logitBias) {
this.options.logitBias = logitBias;
return this;
}
public Builder withLogprobs(Boolean logprobs) {
this.options.logprobs = logprobs;
return this;
}
public Builder withTopLogprobs(Integer topLogprobs) {
this.options.topLogprobs = topLogprobs;
return this;
}
public Builder withMaxTokens(Integer maxTokens) {
this.options.maxTokens = maxTokens;
return this;
}
public Builder withN(Integer n) {
this.options.n = n;
return this;
}
public Builder withPresencePenalty(Float presencePenalty) {
this.options.presencePenalty = presencePenalty;
return this;
}
public Builder withResponseFormat(OpenAiApi.ChatCompletionRequest.ResponseFormat responseFormat) {
this.options.responseFormat = responseFormat;
return this;
}
public Builder withSeed(Integer seed) {
this.options.seed = seed;
return this;
}
public Builder withStop(List<String> stop) {
this.options.stop = stop;
return this;
}
public Builder withTemperature(Float temperature) {
this.options.temperature = temperature;
return this;
}
public Builder withTopP(Float topP) {
this.options.topP = topP;
return this;
}
public Builder withTools(List<OpenAiApi.FunctionTool> tools) {
this.options.tools = tools;
return this;
}
public Builder withToolChoice(String toolChoice) {
this.options.toolChoice = toolChoice;
return this;
}
public Builder withUser(String user) {
this.options.user = user;
return this;
}
public Builder withFunctionCallbacks(List<FunctionCallback> functionCallbacks) {
this.options.functionCallbacks = functionCallbacks;
return this;
}
public Builder withFunctions(Set<String> functionNames) {
Assert.notNull(functionNames, "Function names must not be null");
this.options.functions = functionNames;
return this;
}
public Builder withFunction(String functionName) {
Assert.hasText(functionName, "Function name must not be empty");
this.options.functions.add(functionName);
return this;
}
public FastGPTOptions build() {
return this.options;
}
}
}
package cn.iocoder.yudao.framework.ai.core.model.fastgpt;
import java.util.HashMap;
import org.springframework.ai.chat.metadata.ChatResponseMetadata;
import org.springframework.ai.chat.metadata.EmptyRateLimit;
import org.springframework.ai.chat.metadata.EmptyUsage;
import org.springframework.ai.chat.metadata.RateLimit;
import org.springframework.ai.chat.metadata.Usage;
import org.springframework.ai.openai.metadata.OpenAiRateLimit;
import org.springframework.lang.Nullable;
import org.springframework.util.Assert;
public class FastGPTResponseMetadata extends HashMap<String, Object> implements ChatResponseMetadata {
protected static final String AI_METADATA_STRING = "{ @type: %1$s, id: %2$s, usage: %3$s, rateLimit: %4$s }";
private final String id;
@Nullable
private RateLimit rateLimit;
private final Usage usage;
public static cn.iocoder.yudao.framework.ai.core.model.fastgpt.FastGPTResponseMetadata from(FastGPTApi.ChatCompletion result) {
Assert.notNull(result, "OpenAI ChatCompletionResult must not be null");
FastGPTUsage usage = FastGPTUsage.from(result.usage());
cn.iocoder.yudao.framework.ai.core.model.fastgpt.FastGPTResponseMetadata chatResponseMetadata = new cn.iocoder.yudao.framework.ai.core.model.fastgpt.FastGPTResponseMetadata(result.id(), usage);
return chatResponseMetadata;
}
protected FastGPTResponseMetadata(String id, FastGPTUsage usage) {
this(id, usage, (OpenAiRateLimit)null);
}
protected FastGPTResponseMetadata(String id, FastGPTUsage usage, @Nullable OpenAiRateLimit rateLimit) {
this.id = id;
this.usage = usage;
this.rateLimit = rateLimit;
}
public String getId() {
return this.id;
}
@Nullable
public RateLimit getRateLimit() {
RateLimit rateLimit = this.rateLimit;
return (RateLimit)(rateLimit != null ? rateLimit : new EmptyRateLimit());
}
public Usage getUsage() {
Usage usage = this.usage;
return (Usage)(usage != null ? usage : new EmptyUsage());
}
public cn.iocoder.yudao.framework.ai.core.model.fastgpt.FastGPTResponseMetadata withRateLimit(RateLimit rateLimit) {
this.rateLimit = rateLimit;
return this;
}
public String toString() {
return "{ @type: %1$s, id: %2$s, usage: %3$s, rateLimit: %4$s }".formatted(this.getClass().getName(), this.getId(), this.getUsage(), this.getRateLimit());
}
}
package cn.iocoder.yudao.framework.ai.core.model.fastgpt;
import java.util.ArrayList;
import java.util.List;
import org.springframework.util.CollectionUtils;
public class FastGPTStreamFunctionCallingHelper {
public FastGPTStreamFunctionCallingHelper() {
}
public FastGPTApi.ChatCompletionChunk merge(FastGPTApi.ChatCompletionChunk previous, FastGPTApi.ChatCompletionChunk current) {
if (previous == null) {
return current;
} else {
String id = current.id() != null ? current.id() : previous.id();
Long created = current.created() != null ? current.created() : previous.created();
String model = current.model() != null ? current.model() : previous.model();
String systemFingerprint = current.systemFingerprint() != null ? current.systemFingerprint() : previous.systemFingerprint();
String object = current.object() != null ? current.object() : previous.object();
FastGPTApi.ChatCompletionChunk.ChunkChoice previousChoice0 = CollectionUtils.isEmpty(previous.choices()) ? null : (FastGPTApi.ChatCompletionChunk.ChunkChoice)previous.choices().get(0);
FastGPTApi.ChatCompletionChunk.ChunkChoice currentChoice0 = CollectionUtils.isEmpty(current.choices()) ? null : (FastGPTApi.ChatCompletionChunk.ChunkChoice)current.choices().get(0);
FastGPTApi.ChatCompletionChunk.ChunkChoice choice = this.merge(previousChoice0, currentChoice0);
List<FastGPTApi.ChatCompletionChunk.ChunkChoice> chunkChoices = choice == null ? List.of() : List.of(choice);
return new FastGPTApi.ChatCompletionChunk(id, chunkChoices, created, model, systemFingerprint, object);
}
}
private FastGPTApi.ChatCompletionChunk.ChunkChoice merge(FastGPTApi.ChatCompletionChunk.ChunkChoice previous, FastGPTApi.ChatCompletionChunk.ChunkChoice current) {
if (previous == null) {
return current;
} else {
FastGPTApi.ChatCompletionFinishReason finishReason = current.finishReason() != null ? current.finishReason() : previous.finishReason();
Integer index = current.index() != null ? current.index() : previous.index();
FastGPTApi.ChatCompletionMessage message = this.merge(previous.delta(), current.delta());
FastGPTApi.LogProbs logprobs = current.logprobs() != null ? current.logprobs() : previous.logprobs();
return new FastGPTApi.ChatCompletionChunk.ChunkChoice(finishReason, index, message, logprobs);
}
}
private FastGPTApi.ChatCompletionMessage merge(FastGPTApi.ChatCompletionMessage previous, FastGPTApi.ChatCompletionMessage current) {
String content = current.content() != null ? current.content() : "" + (previous.content() != null ? previous.content() : "");
FastGPTApi.ChatCompletionMessage.Role role = current.role() != null ? current.role() : previous.role();
role = role != null ? role : FastGPTApi.ChatCompletionMessage.Role.ASSISTANT;
String name = current.name() != null ? current.name() : previous.name();
String toolCallId = current.toolCallId() != null ? current.toolCallId() : previous.toolCallId();
List<FastGPTApi.ChatCompletionMessage.ToolCall> toolCalls = new ArrayList();
FastGPTApi.ChatCompletionMessage.ToolCall lastPreviousTooCall = null;
if (previous.toolCalls() != null) {
lastPreviousTooCall = (FastGPTApi.ChatCompletionMessage.ToolCall)previous.toolCalls().get(previous.toolCalls().size() - 1);
if (previous.toolCalls().size() > 1) {
toolCalls.addAll(previous.toolCalls().subList(0, previous.toolCalls().size() - 1));
}
}
if (current.toolCalls() != null) {
if (current.toolCalls().size() > 1) {
throw new IllegalStateException("Currently only one tool call is supported per message!");
}
FastGPTApi.ChatCompletionMessage.ToolCall currentToolCall = (FastGPTApi.ChatCompletionMessage.ToolCall)current.toolCalls().iterator().next();
if (currentToolCall.id() != null) {
if (lastPreviousTooCall != null) {
toolCalls.add(lastPreviousTooCall);
}
toolCalls.add(currentToolCall);
} else {
toolCalls.add(this.merge(lastPreviousTooCall, currentToolCall));
}
} else if (lastPreviousTooCall != null) {
toolCalls.add(lastPreviousTooCall);
}
return new FastGPTApi.ChatCompletionMessage(content, role, name, toolCallId, toolCalls);
}
private FastGPTApi.ChatCompletionMessage.ToolCall merge(FastGPTApi.ChatCompletionMessage.ToolCall previous, FastGPTApi.ChatCompletionMessage.ToolCall current) {
if (previous == null) {
return current;
} else {
String id = current.id() != null ? current.id() : previous.id();
String type = current.type() != null ? current.type() : previous.type();
FastGPTApi.ChatCompletionMessage.ChatCompletionFunction function = this.merge(previous.function(), current.function());
return new FastGPTApi.ChatCompletionMessage.ToolCall(id, type, function);
}
}
private FastGPTApi.ChatCompletionMessage.ChatCompletionFunction merge(FastGPTApi.ChatCompletionMessage.ChatCompletionFunction previous, FastGPTApi.ChatCompletionMessage.ChatCompletionFunction current) {
if (previous == null) {
return current;
} else {
String name = current.name() != null ? current.name() : previous.name();
StringBuilder arguments = new StringBuilder();
if (previous.arguments() != null) {
arguments.append(previous.arguments());
}
if (current.arguments() != null) {
arguments.append(current.arguments());
}
return new FastGPTApi.ChatCompletionMessage.ChatCompletionFunction(name, arguments.toString());
}
}
public boolean isStreamingToolFunctionCall(FastGPTApi.ChatCompletionChunk chatCompletion) {
if (chatCompletion != null && !CollectionUtils.isEmpty(chatCompletion.choices())) {
FastGPTApi.ChatCompletionChunk.ChunkChoice choice = (FastGPTApi.ChatCompletionChunk.ChunkChoice)chatCompletion.choices().get(0);
if (choice != null && choice.delta() != null) {
return !CollectionUtils.isEmpty(choice.delta().toolCalls());
} else {
return false;
}
} else {
return false;
}
}
public boolean isStreamingToolFunctionCallFinish(FastGPTApi.ChatCompletionChunk chatCompletion) {
if (chatCompletion != null && !CollectionUtils.isEmpty(chatCompletion.choices())) {
FastGPTApi.ChatCompletionChunk.ChunkChoice choice = (FastGPTApi.ChatCompletionChunk.ChunkChoice)chatCompletion.choices().get(0);
if (choice != null && choice.delta() != null) {
return choice.finishReason() == FastGPTApi.ChatCompletionFinishReason.TOOL_CALLS;
} else {
return false;
}
} else {
return false;
}
}
public FastGPTApi.ChatCompletion chunkToChatCompletion(FastGPTApi.ChatCompletionChunk chunk) {
List<FastGPTApi.ChatCompletion.Choice> choices = chunk.choices().stream().map((chunkChoice) -> {
return new FastGPTApi.ChatCompletion.Choice(chunkChoice.finishReason(), chunkChoice.index(), chunkChoice.delta(), chunkChoice.logprobs());
}).toList();
return new FastGPTApi.ChatCompletion(chunk.id(), choices, chunk.created(), chunk.model(), chunk.systemFingerprint(), "chat.completion", (FastGPTApi.Usage)null);
}
}
package cn.iocoder.yudao.framework.ai.core.model.fastgpt;
import org.springframework.ai.chat.metadata.Usage;
import org.springframework.util.Assert;
public class FastGPTUsage implements Usage {
private final FastGPTApi.Usage usage;
public static cn.iocoder.yudao.framework.ai.core.model.fastgpt.FastGPTUsage from(FastGPTApi.Usage usage) {
return new cn.iocoder.yudao.framework.ai.core.model.fastgpt.FastGPTUsage(usage);
}
protected FastGPTUsage(FastGPTApi.Usage usage) {
Assert.notNull(usage, "OpenAI Usage must not be null");
this.usage = usage;
}
protected FastGPTApi.Usage getUsage() {
return this.usage;
}
public Long getPromptTokens() {
return this.getUsage().promptTokens().longValue();
}
public Long getGenerationTokens() {
return this.getUsage().completionTokens().longValue();
}
public Long getTotalTokens() {
return this.getUsage().totalTokens().longValue();
}
public String toString() {
return this.getUsage().toString();
}
}
...@@ -3,6 +3,7 @@ package cn.iocoder.yudao.framework.ai.core.util; ...@@ -3,6 +3,7 @@ package cn.iocoder.yudao.framework.ai.core.util;
import cn.hutool.core.util.StrUtil; import cn.hutool.core.util.StrUtil;
import cn.iocoder.yudao.framework.ai.core.enums.AiPlatformEnum; import cn.iocoder.yudao.framework.ai.core.enums.AiPlatformEnum;
import cn.iocoder.yudao.framework.ai.core.model.deepseek.DeepSeekChatOptions; import cn.iocoder.yudao.framework.ai.core.model.deepseek.DeepSeekChatOptions;
import cn.iocoder.yudao.framework.ai.core.model.fastgpt.FastGPTOptions;
import cn.iocoder.yudao.framework.ai.core.model.xinghuo.XingHuoChatOptions; import cn.iocoder.yudao.framework.ai.core.model.xinghuo.XingHuoChatOptions;
import com.alibaba.cloud.ai.tongyi.chat.TongYiChatOptions; import com.alibaba.cloud.ai.tongyi.chat.TongYiChatOptions;
import org.springframework.ai.azure.openai.AzureOpenAiChatOptions; import org.springframework.ai.azure.openai.AzureOpenAiChatOptions;
...@@ -20,7 +21,7 @@ import org.springframework.ai.zhipuai.ZhiPuAiChatOptions; ...@@ -20,7 +21,7 @@ import org.springframework.ai.zhipuai.ZhiPuAiChatOptions;
*/ */
public class AiUtils { public class AiUtils {
public static ChatOptions buildChatOptions(AiPlatformEnum platform, String model, Double temperature, Integer maxTokens) { public static ChatOptions buildChatOptions(AiPlatformEnum platform, String model, Double temperature, Integer maxTokens, String chatId) {
Float temperatureF = temperature != null ? temperature.floatValue() : null; Float temperatureF = temperature != null ? temperature.floatValue() : null;
//noinspection EnhancedSwitchMigration //noinspection EnhancedSwitchMigration
switch (platform) { switch (platform) {
...@@ -41,6 +42,9 @@ public class AiUtils { ...@@ -41,6 +42,9 @@ public class AiUtils {
return AzureOpenAiChatOptions.builder().withDeploymentName(model).withTemperature(temperatureF).withMaxTokens(maxTokens).build(); return AzureOpenAiChatOptions.builder().withDeploymentName(model).withTemperature(temperatureF).withMaxTokens(maxTokens).build();
case OLLAMA: case OLLAMA:
return OllamaOptions.create().withModel(model).withTemperature(temperatureF).withNumPredict(maxTokens); return OllamaOptions.create().withModel(model).withTemperature(temperatureF).withNumPredict(maxTokens);
case FAST_GPT:
return FastGPTOptions.builder().withModel(model).withTemperature(temperatureF).withMaxTokens(maxTokens).withChatId(chatId).build();
// return FastGPTOptions.builder().withModel(model).withTemperature(temperatureF).withMaxTokens(maxTokens).build();
default: default:
throw new IllegalArgumentException(StrUtil.format("未知平台({})", platform)); throw new IllegalArgumentException(StrUtil.format("未知平台({})", platform));
} }
......
...@@ -32,6 +32,9 @@ ...@@ -32,6 +32,9 @@
<el-form-item label="角色设定" prop="systemMessage"> <el-form-item label="角色设定" prop="systemMessage">
<el-input type="textarea" v-model="formData.systemMessage" placeholder="请输入角色设定" /> <el-input type="textarea" v-model="formData.systemMessage" placeholder="请输入角色设定" />
</el-form-item> </el-form-item>
<el-form-item label="AppId" prop="appId">
<el-input v-model="formData.appId" placeholder="请输入FastGPT的AppId" />
</el-form-item>
<el-form-item label="是否公开" prop="publicStatus" v-if="!isUser"> <el-form-item label="是否公开" prop="publicStatus" v-if="!isUser">
<el-radio-group v-model="formData.publicStatus"> <el-radio-group v-model="formData.publicStatus">
<el-radio <el-radio
...@@ -91,7 +94,8 @@ const formData = ref({ ...@@ -91,7 +94,8 @@ const formData = ref({
description: undefined, description: undefined,
systemMessage: undefined, systemMessage: undefined,
publicStatus: true, publicStatus: true,
status: CommonStatusEnum.ENABLE status: CommonStatusEnum.ENABLE,
appId: undefined
}) })
const formRef = ref() // 表单 Ref const formRef = ref() // 表单 Ref
const chatModelList = ref([] as ChatModelVO[]) // 聊天模型列表 const chatModelList = ref([] as ChatModelVO[]) // 聊天模型列表
...@@ -176,7 +180,8 @@ const resetForm = () => { ...@@ -176,7 +180,8 @@ const resetForm = () => {
description: undefined, description: undefined,
systemMessage: undefined, systemMessage: undefined,
publicStatus: true, publicStatus: true,
status: CommonStatusEnum.ENABLE status: CommonStatusEnum.ENABLE,
appId: undefined
} }
formRef.value?.resetFields() formRef.value?.resetFields()
} }
......
...@@ -69,6 +69,7 @@ ...@@ -69,6 +69,7 @@
<el-table-column label="角色类别" align="center" prop="category" /> <el-table-column label="角色类别" align="center" prop="category" />
<el-table-column label="角色描述" align="center" prop="description" /> <el-table-column label="角色描述" align="center" prop="description" />
<el-table-column label="角色设定" align="center" prop="systemMessage" /> <el-table-column label="角色设定" align="center" prop="systemMessage" />
<el-table-column label="AppId" align="center" prop="appId" />
<el-table-column label="是否公开" align="center" prop="publicStatus"> <el-table-column label="是否公开" align="center" prop="publicStatus">
<template #default="scope"> <template #default="scope">
<dict-tag :type="DICT_TYPE.INFRA_BOOLEAN_STRING" :value="scope.row.publicStatus" /> <dict-tag :type="DICT_TYPE.INFRA_BOOLEAN_STRING" :value="scope.row.publicStatus" />
......
Markdown is supported
0% or
You are about to add 0 people to the discussion. Proceed with caution.
Finish editing this message first!
Please register or sign in to comment