From a1d055c8b7b60553c5d579764a3bfb183d7e516d Mon Sep 17 00:00:00 2001 From: "DESKTOP-G2FS5SC\\Administrator" Date: Sat, 27 Jun 2026 13:20:41 +0800 Subject: [PATCH] =?UTF-8?q?LangChain4j=20=E5=88=87=E6=8D=A2=E4=B8=BA=20Spr?= =?UTF-8?q?ing=20AI?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- wemirr-platform-dependencies/pom.xml | 30 +- .../ai-spring-boot-starter/pom.xml | 91 +--- .../ai/autoconfigure/AiAutoConfiguration.java | 9 +- .../embedding/EmbeddingModelFactory.java | 2 +- .../OpenAiEmbeddingModelFactory.java | 31 +- .../embedding/QwenEmbeddingModelFactory.java | 34 +- .../text/DeepSeekTextModelProvider.java | 82 +-- .../text/OpenAiCompatibleChatModel.java | 276 ++++++++++ .../text/OpenAiTextModelProvider.java | 72 +-- .../provider/text/QwenTextModelProvider.java | 54 +- .../core/provider/text/TextModelProvider.java | 4 +- .../core/rag/RetrievalAugmentorBuilder.java | 283 ---------- .../core/rag/TranslationQueryTransformer.java | 115 ----- wemirr-plugin/wemirr-platform-ai/pom.xml | 32 +- .../assistant/interfaces/ChatAssistant.java | 66 +-- .../assistant/service/AssistantService.java | 488 +++++++----------- .../ai/core/model/VectorSearchResult.java | 41 ++ .../ai/core/processor/DocumentProcessor.java | 70 +-- .../processor/VectorizationProcessor.java | 113 ++-- .../embedding/EmbeddingModelService.java | 2 +- .../provider/graph/GraphContentRetriever.java | 121 ----- .../ai/core/provider/graph/GraphDocument.java | 14 + .../ai/core/provider/graph/GraphEdge.java | 12 + .../ai/core/provider/graph/GraphNode.java | 12 + .../core/provider/graph/GraphRagService.java | 89 ++-- .../graph/GraphRagTransformerFactory.java | 308 +++++------ .../core/provider/graph/GraphRetriever.java | 18 - .../ai/core/provider/graph/GraphStore.java | 2 - .../graph/neo4j/Neo4jGraphRetriever.java | 11 +- .../provider/graph/neo4j/Neo4jGraphStore.java | 4 +- .../graph/neo4j/Neo4jGraphWriter.java | 20 +- .../mcp/ContextAwareMcpToolProvider.java | 89 ---- .../provider/mcp/DynamicMcpToolProvider.java | 95 ---- .../ai/core/provider/mcp/McpClientHandle.java | 20 + .../provider/mcp/McpToolProviderFactory.java | 60 --- .../retrieval/ContentRetrieverRegistry.java | 86 --- .../retrieval/ContentRetrieverStrategy.java | 49 -- .../GraphContentRetrieverStrategy.java | 80 --- .../VectorContentRetrieverStrategy.java | 84 --- .../scoring/CohereScoringModelProvider.java | 61 --- .../scoring/JinaScoringModelProvider.java | 54 -- .../scoring/ScoringModelProvider.java | 38 -- .../scoring/ScoringModelProviderRegistry.java | 100 ---- .../provider/scoring/ScoringModelService.java | 91 +--- .../ai/core/provider/text/TextModelCache.java | 4 +- .../core/provider/text/TextModelService.java | 4 +- .../provider/vector/VectorStoreFactory.java | 105 ++-- .../platform/ai/core/sse/SseChatHelper.java | 56 +- .../ai/core/tools/PlatformToolService.java | 8 +- .../agent/impl/AgentNodeExecutor.java | 20 +- .../agent/impl/DocExtractorNodeExecutor.java | 45 +- .../workflow/agent/impl/LLMNodeExecutor.java | 110 ++-- .../impl/ParameterExtractorNodeExecutor.java | 57 +- .../impl/QuestionClassifierNodeExecutor.java | 33 +- .../workflow/agent/impl/ToolNodeExecutor.java | 56 +- .../workflow/runtime/CompiledWorkflow.java | 6 +- .../runtime/LangChain4jWorkflowFactory.java | 152 ------ .../runtime/SpringAiWorkflowFactory.java | 60 +++ .../workflow/runtime/WorkflowCompiler.java | 8 +- .../runtime/WorkflowExecutionScope.java | 2 +- .../runtime/adapter/NodeExecutorAdapter.java | 2 +- .../listener/CustomizeChatModelListener.java | 37 -- .../ai/service/McpConnectionManager.java | 6 +- .../platform/ai/service/ToolService.java | 18 +- .../ai/service/VectorSearchService.java | 5 +- .../ai/service/impl/ChatServiceImpl.java | 9 +- .../impl/ConversationMessageServiceImpl.java | 3 + .../impl/GraphExtractionServiceImpl.java | 9 +- .../ai/service/impl/GraphServiceImpl.java | 9 +- .../impl/KnowledgeSearchServiceImpl.java | 19 +- .../impl/McpConnectionManagerImpl.java | 166 +++--- .../impl/McpServerConfigServiceImpl.java | 13 +- .../service/impl/VectorSearchServiceImpl.java | 50 +- .../impl/WorkflowExecutionServiceImpl.java | 4 +- ...java => SpringAiWorkflowBoundaryTest.java} | 8 +- ....java => SpringAiWorkflowFactoryTest.java} | 24 +- ...rExtractorNodeExecutorIntegrationTest.java | 21 +- ...ClassifierNodeExecutorIntegrationTest.java | 28 +- "\351\231\204\344\273\266/mysql/v4-ai.sql" | 2 +- 79 files changed, 1459 insertions(+), 3113 deletions(-) create mode 100644 wemirr-platform-framework/ai-spring-boot-starter/src/main/java/com/wemirr/framework/ai/core/provider/text/OpenAiCompatibleChatModel.java delete mode 100644 wemirr-platform-framework/ai-spring-boot-starter/src/main/java/com/wemirr/framework/ai/core/rag/RetrievalAugmentorBuilder.java delete mode 100644 wemirr-platform-framework/ai-spring-boot-starter/src/main/java/com/wemirr/framework/ai/core/rag/TranslationQueryTransformer.java create mode 100644 wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/model/VectorSearchResult.java delete mode 100644 wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/graph/GraphContentRetriever.java create mode 100644 wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/graph/GraphDocument.java create mode 100644 wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/graph/GraphEdge.java create mode 100644 wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/graph/GraphNode.java delete mode 100644 wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/mcp/ContextAwareMcpToolProvider.java delete mode 100644 wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/mcp/DynamicMcpToolProvider.java create mode 100644 wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/mcp/McpClientHandle.java delete mode 100644 wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/mcp/McpToolProviderFactory.java delete mode 100644 wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/retrieval/ContentRetrieverRegistry.java delete mode 100644 wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/retrieval/ContentRetrieverStrategy.java delete mode 100644 wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/retrieval/GraphContentRetrieverStrategy.java delete mode 100644 wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/retrieval/VectorContentRetrieverStrategy.java delete mode 100644 wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/scoring/CohereScoringModelProvider.java delete mode 100644 wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/scoring/JinaScoringModelProvider.java delete mode 100644 wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/scoring/ScoringModelProvider.java delete mode 100644 wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/scoring/ScoringModelProviderRegistry.java delete mode 100644 wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/workflow/runtime/LangChain4jWorkflowFactory.java create mode 100644 wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/workflow/runtime/SpringAiWorkflowFactory.java delete mode 100644 wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/listener/CustomizeChatModelListener.java rename wemirr-plugin/wemirr-platform-ai/src/test/java/com/wemirr/platform/ai/core/workflow/runtime/{LangChain4jWorkflowBoundaryTest.java => SpringAiWorkflowBoundaryTest.java} (90%) rename wemirr-plugin/wemirr-platform-ai/src/test/java/com/wemirr/platform/ai/core/workflow/runtime/{LangChain4jWorkflowFactoryTest.java => SpringAiWorkflowFactoryTest.java} (82%) diff --git a/wemirr-platform-dependencies/pom.xml b/wemirr-platform-dependencies/pom.xml index a5e1eaf2..500c4918 100644 --- a/wemirr-platform-dependencies/pom.xml +++ b/wemirr-platform-dependencies/pom.xml @@ -24,7 +24,8 @@ 1.6 3.2.1 2.22.1 - 1.1.2 + 1.1.8 + 1.1.2.3 4.0.2 4.0.3 @@ -56,10 +57,6 @@ 3.1.0 1.45.0 1.8.5-m2 - - 1.9.1 - 1.9.1-beta17 - 1.6.0-beta3 @@ -341,27 +338,18 @@ ${wemirr-platform.version} - + - dev.langchain4j - langchain4j-bom - ${langchain4j.version} + org.springframework.ai + spring-ai-bom + ${spring-ai.version} pom import - dev.langchain4j - langchain4j-community-bom - ${langchain4j.community.version} - pom - import - - - - - org.bsc.langgraph4j - langgraph4j-bom - ${langgraph4j.version} + com.alibaba.cloud.ai + spring-ai-alibaba-bom + ${spring-ai-alibaba.version} pom import diff --git a/wemirr-platform-framework/ai-spring-boot-starter/pom.xml b/wemirr-platform-framework/ai-spring-boot-starter/pom.xml index 3337e699..d6ae7f7d 100644 --- a/wemirr-platform-framework/ai-spring-boot-starter/pom.xml +++ b/wemirr-platform-framework/ai-spring-boot-starter/pom.xml @@ -15,9 +15,8 @@ AI 能力集成 Starter,支持多模型提供商(DeepSeek、Qwen、OpenAI 等) - 1.9.1 - 1.9.1-beta17 - 1.6.0-beta3 + 1.1.8 + 1.1.2.3 @@ -42,70 +41,29 @@ spring-boot-configuration-processor - + - dev.langchain4j - langchain4j + org.springframework.ai + spring-ai-model - dev.langchain4j - langchain4j-core - - - - - - dev.langchain4j - langchain4j-open-ai - - - - dev.langchain4j - langchain4j-community-dashscope - - - - dev.langchain4j - langchain4j-community-qianfan - - - - - dev.langchain4j - langchain4j-jina + org.springframework.ai + spring-ai-openai - dev.langchain4j - langchain4j-cohere + org.springframework.ai + spring-ai-vector-store - - - - dev.langchain4j - langchain4j-web-search-engine-tavily + org.springframework.ai + spring-ai-mcp - + - dev.langchain4j - langchain4j-mcp - - - - - dev.langchain4j - langchain4j-milvus - - - dev.langchain4j - langchain4j-pgvector - - - - - dev.langchain4j - langchain4j-embeddings-all-minilm-l6-v2 + com.alibaba.cloud.ai + spring-ai-alibaba-dashscope + ${spring-ai-alibaba.version} @@ -119,23 +77,16 @@ import - dev.langchain4j - langchain4j-community-bom - ${langchain4j.community.version} - pom - import - - - dev.langchain4j - langchain4j-bom - ${langchain4j.version} + org.springframework.ai + spring-ai-bom + ${spring-ai.version} pom import - org.bsc.langgraph4j - langgraph4j-bom - ${langgraph4j.version} + com.alibaba.cloud.ai + spring-ai-alibaba-bom + ${spring-ai-alibaba.version} pom import diff --git a/wemirr-platform-framework/ai-spring-boot-starter/src/main/java/com/wemirr/framework/ai/autoconfigure/AiAutoConfiguration.java b/wemirr-platform-framework/ai-spring-boot-starter/src/main/java/com/wemirr/framework/ai/autoconfigure/AiAutoConfiguration.java index 69e8c97a..ad8498e5 100644 --- a/wemirr-platform-framework/ai-spring-boot-starter/src/main/java/com/wemirr/framework/ai/autoconfigure/AiAutoConfiguration.java +++ b/wemirr-platform-framework/ai-spring-boot-starter/src/main/java/com/wemirr/framework/ai/autoconfigure/AiAutoConfiguration.java @@ -5,12 +5,12 @@ import com.wemirr.framework.ai.core.provider.embedding.EmbeddingModelRegistry; import com.wemirr.framework.ai.core.provider.embedding.OpenAiEmbeddingModelFactory; import com.wemirr.framework.ai.core.provider.embedding.QwenEmbeddingModelFactory; import com.wemirr.framework.ai.core.provider.text.*; -import dev.langchain4j.model.chat.ChatModel; import org.springframework.boot.autoconfigure.AutoConfiguration; import org.springframework.boot.autoconfigure.condition.ConditionalOnClass; import org.springframework.boot.autoconfigure.condition.ConditionalOnMissingBean; import org.springframework.boot.autoconfigure.condition.ConditionalOnProperty; import org.springframework.boot.context.properties.EnableConfigurationProperties; +import org.springframework.ai.chat.model.ChatModel; import org.springframework.context.annotation.Bean; import java.util.List; @@ -32,7 +32,6 @@ public class AiAutoConfiguration { @Bean @ConditionalOnMissingBean @ConditionalOnProperty(prefix = "ai.provider.deepseek", name = "enabled", havingValue = "true", matchIfMissing = true) - @ConditionalOnClass(name = "dev.langchain4j.model.openai.OpenAiChatModel") public DeepSeekTextModelProvider deepSeekTextModelProvider() { return new DeepSeekTextModelProvider(); } @@ -40,7 +39,6 @@ public class AiAutoConfiguration { @Bean @ConditionalOnMissingBean @ConditionalOnProperty(prefix = "ai.provider.qwen", name = "enabled", havingValue = "true", matchIfMissing = true) - @ConditionalOnClass(name = "dev.langchain4j.community.model.dashscope.QwenChatModel") public QwenTextModelProvider qwenTextModelProvider() { return new QwenTextModelProvider(); } @@ -48,7 +46,6 @@ public class AiAutoConfiguration { @Bean @ConditionalOnMissingBean @ConditionalOnProperty(prefix = "ai.provider.openai", name = "enabled", havingValue = "true", matchIfMissing = true) - @ConditionalOnClass(name = "dev.langchain4j.model.openai.OpenAiChatModel") public OpenAiTextModelProvider openAiTextModelProvider() { return new OpenAiTextModelProvider(); } @@ -64,7 +61,7 @@ public class AiAutoConfiguration { @Bean @ConditionalOnMissingBean @ConditionalOnProperty(prefix = "ai.provider.qwen", name = "enabled", havingValue = "true", matchIfMissing = true) - @ConditionalOnClass(name = "dev.langchain4j.community.model.dashscope.QwenEmbeddingModel") + @ConditionalOnClass(name = "com.alibaba.cloud.ai.dashscope.embedding.text.DashScopeEmbeddingModel") public QwenEmbeddingModelFactory qwenEmbeddingModelFactory() { return new QwenEmbeddingModelFactory(); } @@ -72,7 +69,7 @@ public class AiAutoConfiguration { @Bean @ConditionalOnMissingBean @ConditionalOnProperty(prefix = "ai.provider.openai", name = "enabled", havingValue = "true", matchIfMissing = true) - @ConditionalOnClass(name = "dev.langchain4j.model.openai.OpenAiEmbeddingModel") + @ConditionalOnClass(name = "org.springframework.ai.openai.OpenAiEmbeddingModel") public OpenAiEmbeddingModelFactory openAiEmbeddingModelFactory() { return new OpenAiEmbeddingModelFactory(); } diff --git a/wemirr-platform-framework/ai-spring-boot-starter/src/main/java/com/wemirr/framework/ai/core/provider/embedding/EmbeddingModelFactory.java b/wemirr-platform-framework/ai-spring-boot-starter/src/main/java/com/wemirr/framework/ai/core/provider/embedding/EmbeddingModelFactory.java index 0688d096..c7c2c9e5 100644 --- a/wemirr-platform-framework/ai-spring-boot-starter/src/main/java/com/wemirr/framework/ai/core/provider/embedding/EmbeddingModelFactory.java +++ b/wemirr-platform-framework/ai-spring-boot-starter/src/main/java/com/wemirr/framework/ai/core/provider/embedding/EmbeddingModelFactory.java @@ -2,7 +2,7 @@ package com.wemirr.framework.ai.core.provider.embedding; import com.wemirr.framework.ai.core.enums.AiProvider; import com.wemirr.framework.ai.core.model.ModelConfig; -import dev.langchain4j.model.embedding.EmbeddingModel; +import org.springframework.ai.embedding.EmbeddingModel; /** * 向量模型工厂接口 diff --git a/wemirr-platform-framework/ai-spring-boot-starter/src/main/java/com/wemirr/framework/ai/core/provider/embedding/OpenAiEmbeddingModelFactory.java b/wemirr-platform-framework/ai-spring-boot-starter/src/main/java/com/wemirr/framework/ai/core/provider/embedding/OpenAiEmbeddingModelFactory.java index 07a36230..00f283d9 100644 --- a/wemirr-platform-framework/ai-spring-boot-starter/src/main/java/com/wemirr/framework/ai/core/provider/embedding/OpenAiEmbeddingModelFactory.java +++ b/wemirr-platform-framework/ai-spring-boot-starter/src/main/java/com/wemirr/framework/ai/core/provider/embedding/OpenAiEmbeddingModelFactory.java @@ -2,8 +2,12 @@ package com.wemirr.framework.ai.core.provider.embedding; import com.wemirr.framework.ai.core.enums.AiProvider; import com.wemirr.framework.ai.core.model.ModelConfig; -import dev.langchain4j.model.embedding.EmbeddingModel; -import dev.langchain4j.model.openai.OpenAiEmbeddingModel; +import org.springframework.ai.document.MetadataMode; +import org.springframework.ai.embedding.EmbeddingModel; +import org.springframework.ai.openai.OpenAiEmbeddingModel; +import org.springframework.ai.openai.OpenAiEmbeddingOptions; +import org.springframework.ai.openai.api.OpenAiApi; +import org.springframework.ai.retry.RetryUtils; /** * OpenAI 向量模型工厂 @@ -20,19 +24,30 @@ public class OpenAiEmbeddingModelFactory implements EmbeddingModelFactory { @Override public EmbeddingModel createModel(ModelConfig config) { - OpenAiEmbeddingModel.OpenAiEmbeddingModelBuilder builder = OpenAiEmbeddingModel.builder() - .apiKey(config.getApiKey()) - .modelName(config.getName()); - + OpenAiApi.Builder apiBuilder = OpenAiApi.builder().apiKey(config.getApiKey()); if (config.getBaseUrl() != null && !config.getBaseUrl().isBlank()) { - builder.baseUrl(config.getBaseUrl()); + apiBuilder.baseUrl(config.getBaseUrl()); } - return builder.build(); + OpenAiEmbeddingOptions.Builder optionsBuilder = OpenAiEmbeddingOptions.builder().model(config.getName()); + Integer dimension = extractDimension(config); + if (dimension != null) { + optionsBuilder.dimensions(dimension); + } + return new OpenAiEmbeddingModel(apiBuilder.build(), MetadataMode.NONE, optionsBuilder.build(), + RetryUtils.DEFAULT_RETRY_TEMPLATE); } @Override public AiProvider getProviderType() { return AiProvider.OPEN_AI; } + + protected Integer extractDimension(ModelConfig config) { + if (config.getVariables() == null) { + return null; + } + Object dimension = config.getVariables().get("dimension"); + return dimension instanceof Number number ? number.intValue() : null; + } } diff --git a/wemirr-platform-framework/ai-spring-boot-starter/src/main/java/com/wemirr/framework/ai/core/provider/embedding/QwenEmbeddingModelFactory.java b/wemirr-platform-framework/ai-spring-boot-starter/src/main/java/com/wemirr/framework/ai/core/provider/embedding/QwenEmbeddingModelFactory.java index 0b1ca005..23db40d3 100644 --- a/wemirr-platform-framework/ai-spring-boot-starter/src/main/java/com/wemirr/framework/ai/core/provider/embedding/QwenEmbeddingModelFactory.java +++ b/wemirr-platform-framework/ai-spring-boot-starter/src/main/java/com/wemirr/framework/ai/core/provider/embedding/QwenEmbeddingModelFactory.java @@ -1,9 +1,13 @@ package com.wemirr.framework.ai.core.provider.embedding; +import com.alibaba.cloud.ai.dashscope.api.DashScopeApi; +import com.alibaba.cloud.ai.dashscope.embedding.text.DashScopeEmbeddingModel; +import com.alibaba.cloud.ai.dashscope.embedding.text.DashScopeEmbeddingOptions; import com.wemirr.framework.ai.core.enums.AiProvider; import com.wemirr.framework.ai.core.model.ModelConfig; -import dev.langchain4j.community.model.dashscope.QwenEmbeddingModel; -import dev.langchain4j.model.embedding.EmbeddingModel; +import org.springframework.ai.document.MetadataMode; +import org.springframework.ai.embedding.EmbeddingModel; +import org.springframework.ai.retry.RetryUtils; /** * Qwen 向量模型工厂 @@ -21,9 +25,21 @@ public class QwenEmbeddingModelFactory implements EmbeddingModelFactory { @Override public EmbeddingModel createModel(ModelConfig config) { - return QwenEmbeddingModel.builder() - .apiKey(config.getApiKey()) - .modelName(config.getName()) + DashScopeApi.Builder apiBuilder = DashScopeApi.builder().apiKey(config.getApiKey()); + if (config.getBaseUrl() != null && !config.getBaseUrl().isBlank()) { + apiBuilder.baseUrl(config.getBaseUrl()); + } + DashScopeEmbeddingOptions.Builder optionsBuilder = DashScopeEmbeddingOptions.builder() + .model(config.getName()); + Integer dimension = extractDimension(config); + if (dimension != null) { + optionsBuilder.dimensions(dimension); + } + return DashScopeEmbeddingModel.builder() + .dashScopeApi(apiBuilder.build()) + .metadataMode(MetadataMode.NONE) + .defaultOptions(optionsBuilder.build()) + .retryTemplate(RetryUtils.DEFAULT_RETRY_TEMPLATE) .build(); } @@ -31,4 +47,12 @@ public class QwenEmbeddingModelFactory implements EmbeddingModelFactory { public AiProvider getProviderType() { return AiProvider.QWEN; } + + private Integer extractDimension(ModelConfig config) { + if (config.getVariables() == null) { + return null; + } + Object dimension = config.getVariables().get("dimension"); + return dimension instanceof Number number ? number.intValue() : null; + } } diff --git a/wemirr-platform-framework/ai-spring-boot-starter/src/main/java/com/wemirr/framework/ai/core/provider/text/DeepSeekTextModelProvider.java b/wemirr-platform-framework/ai-spring-boot-starter/src/main/java/com/wemirr/framework/ai/core/provider/text/DeepSeekTextModelProvider.java index 3b2365f5..9c4ec2fb 100644 --- a/wemirr-platform-framework/ai-spring-boot-starter/src/main/java/com/wemirr/framework/ai/core/provider/text/DeepSeekTextModelProvider.java +++ b/wemirr-platform-framework/ai-spring-boot-starter/src/main/java/com/wemirr/framework/ai/core/provider/text/DeepSeekTextModelProvider.java @@ -2,10 +2,6 @@ package com.wemirr.framework.ai.core.provider.text; import com.wemirr.framework.ai.core.enums.AiProvider; import com.wemirr.framework.ai.core.model.ModelConfig; -import dev.langchain4j.model.chat.ChatModel; -import dev.langchain4j.model.chat.StreamingChatModel; -import dev.langchain4j.model.openai.OpenAiChatModel; -import dev.langchain4j.model.openai.OpenAiStreamingChatModel; /** * DeepSeek 文本模型提供者 @@ -13,7 +9,7 @@ import dev.langchain4j.model.openai.OpenAiStreamingChatModel; * @author Levin * @since 2025/10/11 */ -public class DeepSeekTextModelProvider extends AbstractTextModelProvider { +public class DeepSeekTextModelProvider extends OpenAiTextModelProvider { @Override protected AiProvider getAiProvider() { @@ -21,66 +17,20 @@ public class DeepSeekTextModelProvider extends AbstractTextModelProvider { } @Override - public ChatModel createModel(ModelConfig config) { - logModelCreation(config, false); - - OpenAiChatModel.OpenAiChatModelBuilder builder = OpenAiChatModel.builder() - .baseUrl(config.getBaseUrl()) - .apiKey(config.getApiKey()) - .modelName(config.getName()); - - applyCommonParams(builder, config); - - if (isDeepThinkingEnabled(config)) { - builder.returnThinking(true); - } - - return builder.build(); - } - - @Override - public StreamingChatModel createStreamingModel(ModelConfig config) { - logModelCreation(config, true); - - OpenAiStreamingChatModel.OpenAiStreamingChatModelBuilder builder = OpenAiStreamingChatModel.builder() - .baseUrl(config.getBaseUrl()) - .apiKey(config.getApiKey()) - .modelName(config.getName()); - - applyStreamingParams(builder, config); - - if (isDeepThinkingEnabled(config)) { - builder.returnThinking(true); - } - - return builder.build(); - } - - private void applyCommonParams(OpenAiChatModel.OpenAiChatModelBuilder builder, ModelConfig config) { - Integer maxTokens = extractMaxTokens(config); - Double temperature = extractTemperature(config); - Double topP = extractTopP(config); - Double freqPenalty = extractFrequencyPenalty(config); - Double presPenalty = extractPresencePenalty(config); - - if (maxTokens != null) builder.maxTokens(maxTokens); - if (temperature != null) builder.temperature(temperature); - if (topP != null) builder.topP(topP); - if (freqPenalty != null) builder.frequencyPenalty(freqPenalty); - if (presPenalty != null) builder.presencePenalty(presPenalty); - } - - private void applyStreamingParams(OpenAiStreamingChatModel.OpenAiStreamingChatModelBuilder builder, ModelConfig config) { - Integer maxTokens = extractMaxTokens(config); - Double temperature = extractTemperature(config); - Double topP = extractTopP(config); - Double freqPenalty = extractFrequencyPenalty(config); - Double presPenalty = extractPresencePenalty(config); - - if (maxTokens != null) builder.maxTokens(maxTokens); - if (temperature != null) builder.temperature(temperature); - if (topP != null) builder.topP(topP); - if (freqPenalty != null) builder.frequencyPenalty(freqPenalty); - if (presPenalty != null) builder.presencePenalty(presPenalty); + protected OpenAiCompatibleChatModel.Options buildOptions(ModelConfig config) { + OpenAiCompatibleChatModel.Options options = super.buildOptions(config); + return new OpenAiCompatibleChatModel.Options( + options.baseUrl(), + options.apiKey(), + options.model(), + options.maxTokens(), + options.temperature(), + options.topP(), + options.frequencyPenalty(), + options.presencePenalty(), + isDeepThinkingEnabled(config) ? "high" : null, + options.enableSearch(), + options.enableThinking() + ); } } diff --git a/wemirr-platform-framework/ai-spring-boot-starter/src/main/java/com/wemirr/framework/ai/core/provider/text/OpenAiCompatibleChatModel.java b/wemirr-platform-framework/ai-spring-boot-starter/src/main/java/com/wemirr/framework/ai/core/provider/text/OpenAiCompatibleChatModel.java new file mode 100644 index 00000000..bb249c71 --- /dev/null +++ b/wemirr-platform-framework/ai-spring-boot-starter/src/main/java/com/wemirr/framework/ai/core/provider/text/OpenAiCompatibleChatModel.java @@ -0,0 +1,276 @@ +package com.wemirr.framework.ai.core.provider.text; + +import com.fasterxml.jackson.core.JsonProcessingException; +import com.fasterxml.jackson.databind.JsonNode; +import com.fasterxml.jackson.databind.ObjectMapper; +import com.fasterxml.jackson.databind.node.ArrayNode; +import com.fasterxml.jackson.databind.node.ObjectNode; +import lombok.extern.slf4j.Slf4j; +import org.springframework.ai.chat.messages.AssistantMessage; +import org.springframework.ai.chat.messages.Message; +import org.springframework.ai.chat.metadata.ChatGenerationMetadata; +import org.springframework.ai.chat.metadata.ChatResponseMetadata; +import org.springframework.ai.chat.metadata.DefaultUsage; +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.prompt.ChatOptions; +import org.springframework.ai.chat.prompt.Prompt; +import reactor.core.publisher.Flux; + +import java.io.BufferedReader; +import java.io.IOException; +import java.io.InputStream; +import java.io.InputStreamReader; +import java.io.StringReader; +import java.net.URI; +import java.net.http.HttpClient; +import java.net.http.HttpRequest; +import java.net.http.HttpResponse; +import java.nio.charset.StandardCharsets; +import java.time.Duration; +import java.util.List; +import java.util.Locale; +import java.util.Objects; + +/** + * OpenAI 兼容聊天模型。 + *

+ * Spring AI 1.1.8 的 OpenAI HTTP 客户端在 Spring Framework 7 下会调用旧版 Spring Web + * 方法。该实现保留 Spring AI {@link ChatModel} + * 契约,使用 JDK HttpClient 访问 OpenAI-compatible chat completions API。 + *

+ * + * @author Levin + * @since 2026/06/27 + */ +@Slf4j +public class OpenAiCompatibleChatModel implements ChatModel { + + private static final ObjectMapper OBJECT_MAPPER = new ObjectMapper(); + + private static final Duration REQUEST_TIMEOUT = Duration.ofMinutes(3); + + private final HttpClient httpClient = HttpClient.newBuilder() + .connectTimeout(Duration.ofSeconds(30)) + .build(); + + private final Options options; + + public OpenAiCompatibleChatModel(Options options) { + this.options = Objects.requireNonNull(options, "options must not be null"); + } + + @Override + public ChatResponse call(Prompt prompt) { + try { + ObjectNode request = buildRequest(prompt, false); + HttpRequest httpRequest = buildHttpRequest(request); + HttpResponse response = httpClient.send(httpRequest, HttpResponse.BodyHandlers.ofString()); + ensureSuccess(response); + JsonNode root = OBJECT_MAPPER.readTree(response.body()); + JsonNode choice = root.path("choices").path(0); + String content = choice.path("message").path("content").asText(""); + return buildResponse(content, root.path("id").asText(null), + root.path("model").asText(options.model()), choice.path("finish_reason").asText(null), + root.path("usage")); + } catch (IOException e) { + throw new IllegalStateException("调用 OpenAI 兼容聊天模型失败: " + e.getMessage(), e); + } catch (InterruptedException e) { + Thread.currentThread().interrupt(); + throw new IllegalStateException("调用 OpenAI 兼容聊天模型被中断", e); + } + } + + @Override + public Flux stream(Prompt prompt) { + return Flux.create(sink -> { + try { + ObjectNode request = buildRequest(prompt, true); + HttpRequest httpRequest = buildHttpRequest(request); + HttpResponse response = httpClient.send(httpRequest, HttpResponse.BodyHandlers.ofInputStream()); + ensureStreamSuccess(response); + try (BufferedReader reader = new BufferedReader(new InputStreamReader(response.body(), StandardCharsets.UTF_8))) { + String line; + while ((line = reader.readLine()) != null) { + if (sink.isCancelled()) { + return; + } + if (!line.startsWith("data:")) { + continue; + } + String data = line.substring("data:".length()).trim(); + if (data.isBlank() || "[DONE]".equals(data)) { + continue; + } + ChatResponse chatResponse = parseStreamChunk(data); + if (chatResponse != null) { + sink.next(chatResponse); + } + } + } + sink.complete(); + } catch (Exception e) { + sink.error(e); + } + }).doOnCancel(() -> log.debug("OpenAI 兼容聊天流已取消: model={}", options.model())); + } + + @Override + public ChatOptions getDefaultOptions() { + return ChatOptions.builder() + .model(options.model()) + .maxTokens(options.maxTokens()) + .temperature(options.temperature()) + .topP(options.topP()) + .frequencyPenalty(options.frequencyPenalty()) + .presencePenalty(options.presencePenalty()) + .build(); + } + + private ObjectNode buildRequest(Prompt prompt, boolean stream) { + ObjectNode request = OBJECT_MAPPER.createObjectNode(); + request.put("model", options.model()); + request.put("stream", stream); + if (stream) { + ObjectNode streamOptions = request.putObject("stream_options"); + streamOptions.put("include_usage", true); + } + putIfNotNull(request, "max_tokens", options.maxTokens()); + putIfNotNull(request, "temperature", options.temperature()); + putIfNotNull(request, "top_p", options.topP()); + putIfNotNull(request, "frequency_penalty", options.frequencyPenalty()); + putIfNotNull(request, "presence_penalty", options.presencePenalty()); + putIfNotNull(request, "reasoning_effort", options.reasoningEffort()); + putIfNotNull(request, "enable_search", options.enableSearch()); + putIfNotNull(request, "enable_thinking", options.enableThinking()); + + ArrayNode messages = request.putArray("messages"); + List instructions = prompt.getInstructions(); + for (Message message : instructions) { + ObjectNode item = messages.addObject(); + item.put("role", toRole(message)); + item.put("content", message.getText() == null ? "" : message.getText()); + } + return request; + } + + private HttpRequest buildHttpRequest(ObjectNode request) throws JsonProcessingException { + return HttpRequest.newBuilder(chatCompletionsUri()) + .timeout(REQUEST_TIMEOUT) + .header("Authorization", "Bearer " + options.apiKey()) + .header("Content-Type", "application/json") + .header("Accept", "application/json") + .POST(HttpRequest.BodyPublishers.ofString(OBJECT_MAPPER.writeValueAsString(request))) + .build(); + } + + private URI chatCompletionsUri() { + String baseUrl = options.baseUrl(); + if (baseUrl == null || baseUrl.isBlank()) { + baseUrl = "https://api.openai.com/v1"; + } + String normalized = baseUrl.endsWith("/") ? baseUrl.substring(0, baseUrl.length() - 1) : baseUrl; + if (normalized.endsWith("/chat/completions")) { + return URI.create(normalized); + } + return URI.create(normalized + "/chat/completions"); + } + + private void ensureSuccess(HttpResponse response) { + int status = response.statusCode(); + if (status >= 200 && status < 300) { + return; + } + throw new IllegalStateException("OpenAI 兼容聊天模型请求失败: status=" + status + ", body=" + response.body()); + } + + private void ensureStreamSuccess(HttpResponse response) throws IOException { + int status = response.statusCode(); + if (status >= 200 && status < 300) { + return; + } + String body = new String(response.body().readAllBytes(), StandardCharsets.UTF_8); + throw new IllegalStateException("OpenAI 兼容聊天模型流式请求失败: status=" + status + ", body=" + body); + } + + private ChatResponse parseStreamChunk(String data) throws JsonProcessingException { + JsonNode root = OBJECT_MAPPER.readTree(data); + JsonNode choice = root.path("choices").path(0); + JsonNode delta = choice.path("delta"); + String content = delta.path("content").asText(""); + JsonNode usage = root.path("usage"); + String finishReason = choice.path("finish_reason").asText(null); + if (content.isEmpty() && usage.isMissingNode() && finishReason == null) { + return null; + } + return buildResponse(content, root.path("id").asText(null), + root.path("model").asText(options.model()), finishReason, usage); + } + + private ChatResponse buildResponse(String content, String id, String model, String finishReason, JsonNode usageNode) { + ChatGenerationMetadata generationMetadata = ChatGenerationMetadata.builder() + .finishReason(finishReason) + .build(); + Generation generation = new Generation(new AssistantMessage(content), generationMetadata); + ChatResponseMetadata metadata = ChatResponseMetadata.builder() + .id(id) + .model(model) + .usage(toUsage(usageNode)) + .build(); + return new ChatResponse(List.of(generation), metadata); + } + + private DefaultUsage toUsage(JsonNode usageNode) { + if (usageNode == null || usageNode.isMissingNode() || usageNode.isNull()) { + return new DefaultUsage(0, 0, 0); + } + Integer promptTokens = usageNode.path("prompt_tokens").isNumber() ? usageNode.path("prompt_tokens").asInt() : 0; + Integer completionTokens = usageNode.path("completion_tokens").isNumber() ? usageNode.path("completion_tokens").asInt() : 0; + Integer totalTokens = usageNode.path("total_tokens").isNumber() ? usageNode.path("total_tokens").asInt() : promptTokens + completionTokens; + return new DefaultUsage(promptTokens, completionTokens, totalTokens, usageNode); + } + + private String toRole(Message message) { + return message.getMessageType().getValue().toLowerCase(Locale.ROOT); + } + + private void putIfNotNull(ObjectNode node, String field, Integer value) { + if (value != null) { + node.put(field, value); + } + } + + private void putIfNotNull(ObjectNode node, String field, Double value) { + if (value != null) { + node.put(field, value); + } + } + + private void putIfNotNull(ObjectNode node, String field, String value) { + if (value != null && !value.isBlank()) { + node.put(field, value); + } + } + + private void putIfNotNull(ObjectNode node, String field, Boolean value) { + if (value != null) { + node.put(field, value); + } + } + + public record Options( + String baseUrl, + String apiKey, + String model, + Integer maxTokens, + Double temperature, + Double topP, + Double frequencyPenalty, + Double presencePenalty, + String reasoningEffort, + Boolean enableSearch, + Boolean enableThinking + ) { + } +} diff --git a/wemirr-platform-framework/ai-spring-boot-starter/src/main/java/com/wemirr/framework/ai/core/provider/text/OpenAiTextModelProvider.java b/wemirr-platform-framework/ai-spring-boot-starter/src/main/java/com/wemirr/framework/ai/core/provider/text/OpenAiTextModelProvider.java index 5f634a5e..0ccb6ed4 100644 --- a/wemirr-platform-framework/ai-spring-boot-starter/src/main/java/com/wemirr/framework/ai/core/provider/text/OpenAiTextModelProvider.java +++ b/wemirr-platform-framework/ai-spring-boot-starter/src/main/java/com/wemirr/framework/ai/core/provider/text/OpenAiTextModelProvider.java @@ -2,10 +2,8 @@ package com.wemirr.framework.ai.core.provider.text; import com.wemirr.framework.ai.core.enums.AiProvider; import com.wemirr.framework.ai.core.model.ModelConfig; -import dev.langchain4j.model.chat.ChatModel; -import dev.langchain4j.model.chat.StreamingChatModel; -import dev.langchain4j.model.openai.OpenAiChatModel; -import dev.langchain4j.model.openai.OpenAiStreamingChatModel; +import org.springframework.ai.chat.model.ChatModel; +import org.springframework.ai.chat.model.StreamingChatModel; /** * OpenAI 文本模型提供者 @@ -23,62 +21,28 @@ public class OpenAiTextModelProvider extends AbstractTextModelProvider { @Override public ChatModel createModel(ModelConfig config) { logModelCreation(config, false); - - OpenAiChatModel.OpenAiChatModelBuilder builder = OpenAiChatModel.builder() - .apiKey(config.getApiKey()) - .modelName(config.getName()); - - if (config.getBaseUrl() != null && !config.getBaseUrl().isBlank()) { - builder.baseUrl(config.getBaseUrl()); - } - - applyCommonParams(builder, config); - - return builder.build(); + return new OpenAiCompatibleChatModel(buildOptions(config)); } @Override public StreamingChatModel createStreamingModel(ModelConfig config) { logModelCreation(config, true); - - OpenAiStreamingChatModel.OpenAiStreamingChatModelBuilder builder = OpenAiStreamingChatModel.builder() - .apiKey(config.getApiKey()) - .modelName(config.getName()); - - if (config.getBaseUrl() != null && !config.getBaseUrl().isBlank()) { - builder.baseUrl(config.getBaseUrl()); - } - - applyStreamingParams(builder, config); - - return builder.build(); - } - - private void applyCommonParams(OpenAiChatModel.OpenAiChatModelBuilder builder, ModelConfig config) { - Integer maxTokens = extractMaxTokens(config); - Double temperature = extractTemperature(config); - Double topP = extractTopP(config); - Double freqPenalty = extractFrequencyPenalty(config); - Double presPenalty = extractPresencePenalty(config); - - if (maxTokens != null) builder.maxTokens(maxTokens); - if (temperature != null) builder.temperature(temperature); - if (topP != null) builder.topP(topP); - if (freqPenalty != null) builder.frequencyPenalty(freqPenalty); - if (presPenalty != null) builder.presencePenalty(presPenalty); + return createModel(config); } - private void applyStreamingParams(OpenAiStreamingChatModel.OpenAiStreamingChatModelBuilder builder, ModelConfig config) { - Integer maxTokens = extractMaxTokens(config); - Double temperature = extractTemperature(config); - Double topP = extractTopP(config); - Double freqPenalty = extractFrequencyPenalty(config); - Double presPenalty = extractPresencePenalty(config); - - if (maxTokens != null) builder.maxTokens(maxTokens); - if (temperature != null) builder.temperature(temperature); - if (topP != null) builder.topP(topP); - if (freqPenalty != null) builder.frequencyPenalty(freqPenalty); - if (presPenalty != null) builder.presencePenalty(presPenalty); + protected OpenAiCompatibleChatModel.Options buildOptions(ModelConfig config) { + return new OpenAiCompatibleChatModel.Options( + config.getBaseUrl(), + config.getApiKey(), + config.getName(), + extractMaxTokens(config), + extractTemperature(config), + extractTopP(config), + extractFrequencyPenalty(config), + extractPresencePenalty(config), + null, + null, + null + ); } } diff --git a/wemirr-platform-framework/ai-spring-boot-starter/src/main/java/com/wemirr/framework/ai/core/provider/text/QwenTextModelProvider.java b/wemirr-platform-framework/ai-spring-boot-starter/src/main/java/com/wemirr/framework/ai/core/provider/text/QwenTextModelProvider.java index ac5b1e40..b03eb68f 100644 --- a/wemirr-platform-framework/ai-spring-boot-starter/src/main/java/com/wemirr/framework/ai/core/provider/text/QwenTextModelProvider.java +++ b/wemirr-platform-framework/ai-spring-boot-starter/src/main/java/com/wemirr/framework/ai/core/provider/text/QwenTextModelProvider.java @@ -2,10 +2,8 @@ package com.wemirr.framework.ai.core.provider.text; import com.wemirr.framework.ai.core.enums.AiProvider; import com.wemirr.framework.ai.core.model.ModelConfig; -import dev.langchain4j.community.model.dashscope.QwenChatModel; -import dev.langchain4j.community.model.dashscope.QwenStreamingChatModel; -import dev.langchain4j.model.chat.ChatModel; -import dev.langchain4j.model.chat.StreamingChatModel; +import org.springframework.ai.chat.model.ChatModel; +import org.springframework.ai.chat.model.StreamingChatModel; /** * 通义千问文本模型提供者 @@ -15,6 +13,8 @@ import dev.langchain4j.model.chat.StreamingChatModel; */ public class QwenTextModelProvider extends AbstractTextModelProvider { + private static final String DEFAULT_DASHSCOPE_COMPATIBLE_BASE_URL = "https://dashscope.aliyuncs.com/compatible-mode/v1"; + @Override protected AiProvider getAiProvider() { return AiProvider.QWEN; @@ -23,40 +23,26 @@ public class QwenTextModelProvider extends AbstractTextModelProvider { @Override public ChatModel createModel(ModelConfig config) { logModelCreation(config, false); - - QwenChatModel.QwenChatModelBuilder builder = QwenChatModel.builder() - .apiKey(config.getApiKey()) - .modelName(config.getName()); - - Integer maxTokens = extractMaxTokens(config); - if (maxTokens != null) { - builder.maxTokens(maxTokens); - } - - if (isWebSearchEnabled(config)) { - builder.enableSearch(true); - } - - return builder.build(); + String baseUrl = config.getBaseUrl() == null || config.getBaseUrl().isBlank() + ? DEFAULT_DASHSCOPE_COMPATIBLE_BASE_URL : config.getBaseUrl(); + return new OpenAiCompatibleChatModel(new OpenAiCompatibleChatModel.Options( + baseUrl, + config.getApiKey(), + config.getName(), + extractMaxTokens(config), + extractTemperature(config), + extractTopP(config), + extractFrequencyPenalty(config), + extractPresencePenalty(config), + null, + isWebSearchEnabled(config) ? Boolean.TRUE : null, + isDeepThinkingEnabled(config) ? Boolean.TRUE : null + )); } @Override public StreamingChatModel createStreamingModel(ModelConfig config) { logModelCreation(config, true); - - QwenStreamingChatModel.QwenStreamingChatModelBuilder builder = QwenStreamingChatModel.builder() - .apiKey(config.getApiKey()) - .modelName(config.getName()); - - Integer maxTokens = extractMaxTokens(config); - if (maxTokens != null) { - builder.maxTokens(maxTokens); - } - - if (isWebSearchEnabled(config)) { - builder.enableSearch(true); - } - - return builder.build(); + return createModel(config); } } diff --git a/wemirr-platform-framework/ai-spring-boot-starter/src/main/java/com/wemirr/framework/ai/core/provider/text/TextModelProvider.java b/wemirr-platform-framework/ai-spring-boot-starter/src/main/java/com/wemirr/framework/ai/core/provider/text/TextModelProvider.java index 9fefa988..86ebbd32 100644 --- a/wemirr-platform-framework/ai-spring-boot-starter/src/main/java/com/wemirr/framework/ai/core/provider/text/TextModelProvider.java +++ b/wemirr-platform-framework/ai-spring-boot-starter/src/main/java/com/wemirr/framework/ai/core/provider/text/TextModelProvider.java @@ -3,8 +3,8 @@ package com.wemirr.framework.ai.core.provider.text; import com.wemirr.framework.ai.core.enums.ModelType; import com.wemirr.framework.ai.core.model.ModelConfig; import com.wemirr.framework.ai.core.provider.ModelProvider; -import dev.langchain4j.model.chat.ChatModel; -import dev.langchain4j.model.chat.StreamingChatModel; +import org.springframework.ai.chat.model.ChatModel; +import org.springframework.ai.chat.model.StreamingChatModel; /** * 文本模型提供者接口 diff --git a/wemirr-platform-framework/ai-spring-boot-starter/src/main/java/com/wemirr/framework/ai/core/rag/RetrievalAugmentorBuilder.java b/wemirr-platform-framework/ai-spring-boot-starter/src/main/java/com/wemirr/framework/ai/core/rag/RetrievalAugmentorBuilder.java deleted file mode 100644 index d81dcd22..00000000 --- a/wemirr-platform-framework/ai-spring-boot-starter/src/main/java/com/wemirr/framework/ai/core/rag/RetrievalAugmentorBuilder.java +++ /dev/null @@ -1,283 +0,0 @@ -package com.wemirr.framework.ai.core.rag; - -import dev.langchain4j.data.segment.TextSegment; -import dev.langchain4j.model.chat.ChatModel; -import dev.langchain4j.model.embedding.EmbeddingModel; -import dev.langchain4j.model.scoring.ScoringModel; -import dev.langchain4j.rag.DefaultRetrievalAugmentor; -import dev.langchain4j.rag.RetrievalAugmentor; -import dev.langchain4j.rag.content.aggregator.ContentAggregator; -import dev.langchain4j.rag.content.aggregator.DefaultContentAggregator; -import dev.langchain4j.rag.content.aggregator.ReRankingContentAggregator; -import dev.langchain4j.rag.content.injector.ContentInjector; -import dev.langchain4j.rag.content.injector.DefaultContentInjector; -import dev.langchain4j.rag.content.retriever.ContentRetriever; -import dev.langchain4j.rag.content.retriever.EmbeddingStoreContentRetriever; -import dev.langchain4j.rag.content.retriever.WebSearchContentRetriever; -import dev.langchain4j.rag.query.router.DefaultQueryRouter; -import dev.langchain4j.rag.query.router.QueryRouter; -import dev.langchain4j.rag.query.transformer.CompressingQueryTransformer; -import dev.langchain4j.rag.query.transformer.QueryTransformer; -import dev.langchain4j.store.embedding.EmbeddingStore; -import dev.langchain4j.web.search.WebSearchEngine; -import dev.langchain4j.web.search.tavily.TavilyWebSearchEngine; -import org.slf4j.Logger; -import org.slf4j.LoggerFactory; - -import java.util.ArrayList; -import java.util.LinkedHashMap; -import java.util.List; -import java.util.Map; -import java.util.concurrent.ExecutorService; -import java.util.concurrent.Executors; - -/** - * RAG 检索增强器构建器 - *

- * 支持多种检索模式:向量检索、Web搜索,可灵活组合使用 - *

- * - * @author Levin - * @since 2025/10/22 - */ -public class RetrievalAugmentorBuilder { - - private static final Logger log = LoggerFactory.getLogger(RetrievalAugmentorBuilder.class); - - private ChatModel chatModel; - private EmbeddingModel embeddingModel; - private EmbeddingStore embeddingStore; - - // 向量检索配置 - private boolean enableVectorRetrieval = true; - private int embeddingMaxResults = 8; - - // Web 搜索配置 - private boolean enableWebSearch = false; - private String webSearchApiKey; - private String webSearchEngineType = "tavily"; - private int webMaxResults = 5; - - // 增强策略 - private boolean enableQueryCompression = true; - private boolean enableReRanking = false; - private boolean enableParallelRetrieval = true; - private ExecutorService executorService; - - // 重排序配置 - private ScoringModel scoringModel; - private int rerankMaxResults = 5; - private double rerankMinScore = 0.5; - - // 自定义检索器 - private final List customRetrievers = new ArrayList<>(); - - private RetrievalAugmentorBuilder() {} - - public static RetrievalAugmentorBuilder builder() { - return new RetrievalAugmentorBuilder(); - } - - // === 配置方法 === - - public RetrievalAugmentorBuilder chatModel(ChatModel chatModel) { - this.chatModel = chatModel; - return this; - } - - public RetrievalAugmentorBuilder embeddingModel(EmbeddingModel embeddingModel) { - this.embeddingModel = embeddingModel; - return this; - } - - public RetrievalAugmentorBuilder embeddingStore(EmbeddingStore embeddingStore) { - this.embeddingStore = embeddingStore; - return this; - } - - public RetrievalAugmentorBuilder enableVectorRetrieval(boolean enable) { - this.enableVectorRetrieval = enable; - return this; - } - - public RetrievalAugmentorBuilder enableWebSearch(boolean enable) { - this.enableWebSearch = enable; - return this; - } - - public RetrievalAugmentorBuilder webSearchApiKey(String apiKey) { - this.webSearchApiKey = apiKey; - return this; - } - - public RetrievalAugmentorBuilder webSearchEngineType(String type) { - this.webSearchEngineType = type; - return this; - } - - public RetrievalAugmentorBuilder webMaxResults(int maxResults) { - this.webMaxResults = maxResults; - return this; - } - - public RetrievalAugmentorBuilder embeddingMaxResults(int maxResults) { - this.embeddingMaxResults = maxResults; - return this; - } - - public RetrievalAugmentorBuilder enableQueryCompression(boolean enable) { - this.enableQueryCompression = enable; - return this; - } - - public RetrievalAugmentorBuilder enableReRanking(boolean enable) { - this.enableReRanking = enable; - return this; - } - - public RetrievalAugmentorBuilder scoringModel(ScoringModel scoringModel) { - this.scoringModel = scoringModel; - return this; - } - - public RetrievalAugmentorBuilder rerankMaxResults(int maxResults) { - this.rerankMaxResults = maxResults; - return this; - } - - public RetrievalAugmentorBuilder rerankMinScore(double minScore) { - this.rerankMinScore = minScore; - return this; - } - - public RetrievalAugmentorBuilder enableParallelRetrieval(boolean enable) { - this.enableParallelRetrieval = enable; - return this; - } - - public RetrievalAugmentorBuilder executorService(ExecutorService executor) { - this.executorService = executor; - return this; - } - - /** - * 添加自定义检索器(如图谱检索器) - */ - public RetrievalAugmentorBuilder addRetriever(ContentRetriever retriever) { - this.customRetrievers.add(retriever); - return this; - } - - // === 构建方法 === - - public RetrievalAugmentor build() { - validate(); - - QueryTransformer queryTransformer = null; - if (enableQueryCompression) { - queryTransformer = new CompressingQueryTransformer(chatModel); - log.debug("Using CompressingQueryTransformer"); - } - - Map retrieverToDescription = new LinkedHashMap<>(); - - // 1. 向量库检索器 - if (enableVectorRetrieval && embeddingStore != null && embeddingModel != null) { - EmbeddingStoreContentRetriever vectorRetriever = EmbeddingStoreContentRetriever.builder() - .embeddingStore(embeddingStore) - .embeddingModel(embeddingModel) - .maxResults(embeddingMaxResults) - .build(); - retrieverToDescription.put(vectorRetriever, "Internal knowledge base"); - log.debug("Enabled vector retrieval with maxResults={}", embeddingMaxResults); - } - - // 2. Web 搜索检索器 - if (enableWebSearch && webSearchApiKey != null && !webSearchApiKey.trim().isEmpty()) { - WebSearchEngine webSearchEngine = createWebSearchEngine(); - WebSearchContentRetriever webRetriever = WebSearchContentRetriever.builder() - .webSearchEngine(webSearchEngine) - .maxResults(webMaxResults) - .build(); - retrieverToDescription.put(webRetriever, "Real-time web information"); - log.debug("Enabled web search ({}) with maxResults={}", webSearchEngineType, webMaxResults); - } - - // 3. 自定义检索器 - for (ContentRetriever customRetriever : customRetrievers) { - retrieverToDescription.put(customRetriever, "Custom retriever"); - log.debug("Added custom retriever: {}", customRetriever.getClass().getSimpleName()); - } - - List retrievers = new ArrayList<>(retrieverToDescription.keySet()); - - if (retrievers.isEmpty()) { - throw new IllegalStateException("No retriever is enabled. Please enable at least one."); - } - - QueryRouter queryRouter = new DefaultQueryRouter(retrievers.toArray(new ContentRetriever[0])); - log.debug("Enabled multi-retrieval with {} retrievers", retrievers.size()); - - ContentAggregator contentAggregator = enableReRanking ? - createReRankingAggregator() : - new DefaultContentAggregator(); - - ContentInjector contentInjector = new DefaultContentInjector(); - - ExecutorService finalExecutor = this.executorService; - if (enableParallelRetrieval && finalExecutor == null) { - finalExecutor = Executors.newFixedThreadPool(Math.min(retrieverToDescription.size(), 4)); - } - - DefaultRetrievalAugmentor.DefaultRetrievalAugmentorBuilder builder = DefaultRetrievalAugmentor.builder() - .queryTransformer(queryTransformer) - .queryRouter(queryRouter) - .contentAggregator(contentAggregator) - .contentInjector(contentInjector); - - if (finalExecutor != null) { - builder.executor(finalExecutor); - } - - return builder.build(); - } - - private WebSearchEngine createWebSearchEngine() { - return switch (webSearchEngineType.toLowerCase()) { - case "tavily" -> TavilyWebSearchEngine.builder().apiKey(webSearchApiKey).build(); - default -> throw new IllegalArgumentException("Unsupported web search engine: " + webSearchEngineType); - }; - } - - private ContentAggregator createReRankingAggregator() { - if (scoringModel == null) { - log.warn("重排序已启用但未配置 ScoringModel,降级使用默认聚合器"); - return new DefaultContentAggregator(); - } - - log.debug("启用重排序: maxResults={}, minScore={}", rerankMaxResults, rerankMinScore); - return ReRankingContentAggregator.builder() - .scoringModel(scoringModel) - .maxResults(rerankMaxResults) - .minScore(rerankMinScore) - .build(); - } - - private void validate() { - if (chatModel == null) { - throw new IllegalArgumentException("chatModel is required"); - } - if (enableVectorRetrieval && (embeddingStore == null || embeddingModel == null)) { - throw new IllegalArgumentException("embeddingStore and embeddingModel are required when vector retrieval is enabled"); - } - if (enableWebSearch && (webSearchApiKey == null || webSearchApiKey.trim().isEmpty())) { - throw new IllegalArgumentException("webSearchApiKey is required when web search is enabled"); - } - } - - public void shutdown() { - if (executorService != null && !executorService.isShutdown()) { - executorService.shutdown(); - } - } -} diff --git a/wemirr-platform-framework/ai-spring-boot-starter/src/main/java/com/wemirr/framework/ai/core/rag/TranslationQueryTransformer.java b/wemirr-platform-framework/ai-spring-boot-starter/src/main/java/com/wemirr/framework/ai/core/rag/TranslationQueryTransformer.java deleted file mode 100644 index 3ba3544a..00000000 --- a/wemirr-platform-framework/ai-spring-boot-starter/src/main/java/com/wemirr/framework/ai/core/rag/TranslationQueryTransformer.java +++ /dev/null @@ -1,115 +0,0 @@ -package com.wemirr.framework.ai.core.rag; - -import dev.langchain4j.model.chat.ChatModel; -import dev.langchain4j.model.input.Prompt; -import dev.langchain4j.model.input.PromptTemplate; -import dev.langchain4j.rag.query.Query; -import dev.langchain4j.rag.query.transformer.QueryTransformer; - -import java.util.Collection; -import java.util.HashMap; -import java.util.Map; - -import static dev.langchain4j.internal.Utils.getOrDefault; -import static dev.langchain4j.internal.ValidationUtils.ensureNotNull; -import static java.util.Collections.singletonList; - -/** - * 将用户查询翻译为目标语言(如中文),以便在单语知识库中进行检索 - * - * @author Levin - * @since 2025/11/3 - */ -public class TranslationQueryTransformer implements QueryTransformer { - - public static final PromptTemplate DEFAULT_PROMPT_TEMPLATE = PromptTemplate.from( - """ - You are a translation assistant. Your task is to translate the user's query into Chinese for information retrieval from a Chinese knowledge base. - - Important rules: - 1. Translate the user's query accurately into Chinese - 2. Preserve all technical terms, proper nouns, and key concepts - 3. Maintain the original intent and meaning - 4. Keep the query concise and suitable for search - 5. If the query is already in Chinese, return it as-is with minor improvements if needed - 6. Only return the translated query, nothing else! - - User query: {{query}} - - Translated Chinese query:""" - ); - - protected final PromptTemplate promptTemplate; - protected final ChatModel chatLanguageModel; - - public TranslationQueryTransformer(ChatModel chatLanguageModel) { - this(chatLanguageModel, DEFAULT_PROMPT_TEMPLATE); - } - - public TranslationQueryTransformer(ChatModel chatLanguageModel, PromptTemplate promptTemplate) { - this.chatLanguageModel = ensureNotNull(chatLanguageModel, "chatLanguageModel"); - this.promptTemplate = getOrDefault(promptTemplate, DEFAULT_PROMPT_TEMPLATE); - } - - public static Builder builder() { - return new Builder(); - } - - @Override - public Collection transform(Query query) { - String originalQuery = query.text(); - - if (isLikelyChinese(originalQuery)) { - return singletonList(query); - } - - Prompt prompt = createPrompt(query); - String translatedQueryText = chatLanguageModel.chat(prompt.text()); - translatedQueryText = cleanTranslatedText(translatedQueryText); - - Query translatedQuery = query.metadata() == null - ? Query.from(translatedQueryText) - : Query.from(translatedQueryText, query.metadata()); - return singletonList(translatedQuery); - } - - protected Prompt createPrompt(Query query) { - Map variables = new HashMap<>(); - variables.put("query", query.text()); - return promptTemplate.apply(variables); - } - - private boolean isLikelyChinese(String text) { - if (text == null || text.trim().isEmpty()) { - return false; - } - return text.chars().anyMatch(ch -> - Character.UnicodeScript.of(ch) == Character.UnicodeScript.HAN); - } - - private String cleanTranslatedText(String text) { - if (text == null) { - return ""; - } - return text.replaceAll("^(翻译后的查询|Translated query|中文查询):\\s*", "").trim(); - } - - public static class Builder { - private ChatModel chatLanguageModel; - private PromptTemplate promptTemplate; - - public Builder chatLanguageModel(ChatModel chatLanguageModel) { - this.chatLanguageModel = chatLanguageModel; - return this; - } - - public Builder promptTemplate(PromptTemplate promptTemplate) { - this.promptTemplate = promptTemplate; - return this; - } - - public TranslationQueryTransformer build() { - return new TranslationQueryTransformer(this.chatLanguageModel, this.promptTemplate); - } - } -} diff --git a/wemirr-plugin/wemirr-platform-ai/pom.xml b/wemirr-plugin/wemirr-platform-ai/pom.xml index c3b766a4..337e70f4 100644 --- a/wemirr-plugin/wemirr-platform-ai/pom.xml +++ b/wemirr-plugin/wemirr-platform-ai/pom.xml @@ -17,6 +17,7 @@ 21 UTF-8 true + 1.1.8 @@ -75,25 +76,24 @@ 2.6.4 - - + - dev.langchain4j - langchain4j-document-parser-apache-tika + org.springframework.ai + spring-ai-tika-document-reader + ${spring-ai.version} - - - - - dev.langchain4j - langchain4j-community-llm-graph-transformer + org.springframework.ai + spring-ai-milvus-store + ${spring-ai.version} - - dev.langchain4j - langchain4j-community-neo4j + org.springframework.ai + spring-ai-pgvector-store + ${spring-ai.version} + + org.neo4j.driver @@ -111,12 +111,6 @@ jackson-databind - - dev.langchain4j - langchain4j-agentic - 1.10.0-beta18 - - diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/assistant/interfaces/ChatAssistant.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/assistant/interfaces/ChatAssistant.java index af3bd2e1..aeefd156 100644 --- a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/assistant/interfaces/ChatAssistant.java +++ b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/assistant/interfaces/ChatAssistant.java @@ -1,50 +1,36 @@ -package com.wemirr.platform.ai.core.assistant.interfaces; +/* + * Copyright (c) 2023 WEMIRR-PLATFORM Authors. All Rights Reserved. + * + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ -import dev.langchain4j.data.message.ImageContent; -import dev.langchain4j.memory.ChatMemory; -import dev.langchain4j.service.*; -import dev.langchain4j.service.memory.ChatMemoryAccess; +package com.wemirr.platform.ai.core.assistant.interfaces; -import java.util.List; +import org.springframework.ai.chat.model.ChatResponse; +import reactor.core.publisher.Flux; /** + * Spring AI 聊天助手抽象。 + * * @author xJh * @since 2025/10/11 **/ -public interface ChatAssistant extends ChatMemoryAccess { - - /** - * 流式聊天助手 - */ - TokenStream chatStream(@MemoryId Long memoryId, @UserMessage String userMessage); - - /** - * 普通rag问答 - */ - TokenStream chatRag(@MemoryId Long memoryId, @UserMessage String message); - - /** - * 聊天助手,附带系统消息 - * @param memoryId - * @param systemMessage - * @param prompt - * @param images - * @return - */ - @SystemMessage("{{sm}}") - TokenStream chatWithSystem(@MemoryId Long memoryId, @V("sm") String systemMessage, @UserMessage String prompt, @UserMessage List images); - - /** - * 聊天助手,不带系统消息 - * @param memoryId - * @param prompt - * @param images - * @return - */ - TokenStream chat(@MemoryId Long memoryId, @UserMessage String prompt, @UserMessage List images); +public interface ChatAssistant { - // 提供访问和清除特定用户记忆的方法,增强可控性 - boolean evictChatMemory(@MemoryId int memoryId); - ChatMemory getChatMemory(@MemoryId int memoryId); + Flux chatStream(Long memoryId, String userMessage); + ChatResponse chat(Long memoryId, String userMessage); } diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/assistant/service/AssistantService.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/assistant/service/AssistantService.java index 8e2e5d5f..d18838ac 100644 --- a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/assistant/service/AssistantService.java +++ b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/assistant/service/AssistantService.java @@ -1,371 +1,263 @@ package com.wemirr.platform.ai.core.assistant.service; import cn.hutool.core.collection.CollUtil; -import com.fasterxml.jackson.databind.ObjectMapper; -import com.wemirr.framework.ai.core.provider.embedding.EmbeddingModelRegistry; -import com.wemirr.framework.ai.core.rag.TranslationQueryTransformer; +import cn.hutool.core.util.StrUtil; import com.wemirr.framework.commons.exception.CheckedException; +import com.wemirr.framework.db.mybatisplus.wrap.Wraps; import com.wemirr.platform.ai.core.assistant.interfaces.ChatAssistant; -import com.wemirr.platform.ai.core.provider.graph.GraphContentRetriever; +import com.wemirr.platform.ai.core.constant.AiServiceConstants; +import com.wemirr.platform.ai.core.enums.MessageRole; +import com.wemirr.platform.ai.core.model.VectorSearchResult; import com.wemirr.platform.ai.core.provider.graph.GraphRagService; -import com.wemirr.platform.ai.core.provider.mcp.McpToolProviderFactory; -import com.wemirr.platform.ai.core.provider.scoring.ScoringModelService; import com.wemirr.platform.ai.core.provider.text.TextModelService; -import com.wemirr.platform.ai.core.provider.vector.VectorStoreFactory; import com.wemirr.platform.ai.domain.entity.ChatAgent; +import com.wemirr.platform.ai.domain.entity.ConversationTurn; import com.wemirr.platform.ai.domain.entity.KnowledgeBase; import com.wemirr.platform.ai.domain.entity.ModelEntity; +import com.wemirr.platform.ai.service.ConversationMessageService; import com.wemirr.platform.ai.service.KnowledgeBaseService; import com.wemirr.platform.ai.service.ToolService; -import dev.langchain4j.data.segment.TextSegment; -import dev.langchain4j.memory.chat.ChatMemoryProvider; -import dev.langchain4j.memory.chat.MessageWindowChatMemory; -import dev.langchain4j.model.chat.ChatModel; -import dev.langchain4j.model.chat.StreamingChatModel; -import dev.langchain4j.model.embedding.EmbeddingModel; -import dev.langchain4j.model.scoring.ScoringModel; -import dev.langchain4j.rag.DefaultRetrievalAugmentor; -import dev.langchain4j.rag.RetrievalAugmentor; -import dev.langchain4j.rag.content.aggregator.ContentAggregator; -import dev.langchain4j.rag.content.aggregator.DefaultContentAggregator; -import dev.langchain4j.rag.content.aggregator.ReRankingContentAggregator; -import dev.langchain4j.rag.content.injector.ContentInjector; -import dev.langchain4j.rag.content.injector.DefaultContentInjector; -import dev.langchain4j.rag.content.retriever.ContentRetriever; -import dev.langchain4j.rag.content.retriever.EmbeddingStoreContentRetriever; -import dev.langchain4j.rag.query.Query; -import dev.langchain4j.rag.query.router.DefaultQueryRouter; -import dev.langchain4j.rag.query.router.QueryRouter; -import dev.langchain4j.rag.query.transformer.CompressingQueryTransformer; -import dev.langchain4j.rag.query.transformer.QueryTransformer; -import dev.langchain4j.service.AiServices; -import dev.langchain4j.service.tool.ToolProvider; -import dev.langchain4j.store.embedding.EmbeddingStore; +import com.wemirr.platform.ai.service.VectorSearchService; import lombok.RequiredArgsConstructor; -import lombok.SneakyThrows; import lombok.extern.slf4j.Slf4j; -import org.apache.commons.lang3.StringUtils; +import org.springframework.ai.chat.messages.AssistantMessage; +import org.springframework.ai.chat.messages.Message; +import org.springframework.ai.chat.messages.SystemMessage; +import org.springframework.ai.chat.messages.UserMessage; +import org.springframework.ai.chat.model.ChatModel; +import org.springframework.ai.chat.model.ChatResponse; +import org.springframework.ai.chat.prompt.Prompt; import org.springframework.beans.factory.annotation.Autowired; -import org.springframework.context.ApplicationContext; import org.springframework.stereotype.Service; - -import java.util.*; -import java.util.concurrent.Executors; +import reactor.core.publisher.Flux; + +import java.util.ArrayList; +import java.util.Collections; +import java.util.Comparator; +import java.util.LinkedHashSet; +import java.util.List; +import java.util.Objects; +import java.util.Set; import java.util.function.Function; import java.util.stream.Collectors; import static com.wemirr.platform.ai.core.constant.AiServiceConstants.DEFAULT_MAX_MESSAGES; /** - * AI助手服务 - *

- * 负责创建不同类型的AI助手: - *

    - *
  • 普通记忆对话助手
  • - *
  • RAG知识库对话助手
  • - *
  • 智能体助手(支持Tools和RAG)
  • - *
+ * AI 助手服务。 * * @author xJh * @since 2025/10/11 - * @apiNote Langchain4j暂未集成稀疏向量用于多路检索 */ @Slf4j @Service @RequiredArgsConstructor public class AssistantService { + private static final String DEFAULT_SYSTEM_PROMPT = "你是一个专业、可靠的企业级智能助手。"; + private final TextModelService textModelService; - private final VectorStoreFactory vectorStoreFactory; private final KnowledgeBaseService knowledgeBaseService; - private final EmbeddingModelRegistry embeddingModelRegistry; - private final ApplicationContext applicationContext; - private final ObjectMapper objectMapper; + private final ConversationMessageService conversationMessageService; + private final VectorSearchService vectorSearchService; private final ToolService toolService; - private final McpToolProviderFactory mcpToolProviderFactory; - private final ScoringModelService scoringModelService; - /** - * GraphRAG 服务(可选,仅在启用图谱功能时注入) - */ @Autowired(required = false) private GraphRagService graphRagService; - /** - * 创建普通记忆对话的 Assistant - * @param modelEntity 模型配置 - * @return ChatAssistant 实例 - */ public ChatAssistant createMemoryAssistant(ModelEntity modelEntity) { - ChatModel chatModel = textModelService.model(modelEntity); - StreamingChatModel streamModel = textModelService.streamModel(modelEntity); - - return AiServices.builder(ChatAssistant.class) - .chatModel(chatModel) - .streamingChatModel(streamModel) - .chatMemory(MessageWindowChatMemory.withMaxMessages(DEFAULT_MAX_MESSAGES)) - .chatMemoryProvider(createMemoryProvider()) - .build(); + return createAssistant(chatModel, message -> DEFAULT_SYSTEM_PROMPT, DEFAULT_MAX_MESSAGES); } - - /** - * 创建智能体对话助手 (支持Tools和RAG、MCP工具) - */ - @SneakyThrows public ChatAssistant createAgentAssistant(ChatAgent chatAgent, ModelEntity modelEntity, RagAssistantParams ragParams) { ChatModel chatModel = textModelService.model(modelEntity); - StreamingChatModel streamModel = textModelService.streamModel(modelEntity); - - var builder = AiServices.builder(ChatAssistant.class) - .chatModel(chatModel) - .streamingChatModel(streamModel) - .chatMemory(MessageWindowChatMemory.withMaxMessages(DEFAULT_MAX_MESSAGES)) - .chatMemoryProvider(createMemoryProvider()); + return createAssistant(chatModel, message -> buildAgentSystemPrompt(chatAgent, ragParams, message), DEFAULT_MAX_MESSAGES); + } - // Configure Tools - if (CollUtil.isNotEmpty(chatAgent.getTools())) { - List tools = chatAgent.getTools().stream() - .map(name -> { - try { - return applicationContext.getBean(name); - } catch (Exception e) { - log.warn("Tool bean not found: {}", name); - return null; - } - }) - .filter(Objects::nonNull) - .toList(); - if (!tools.isEmpty()) { - builder.tools(tools); - } - } + public ChatAssistant createMemoryRagAssistant(RagAssistantParams params) { + ChatModel chatModel = textModelService.model(params.getTextModelEntity()); + int maxMessages = params.getMaxMessages() == null ? DEFAULT_MAX_MESSAGES : params.getMaxMessages(); + return createAssistant(chatModel, message -> buildRagSystemPrompt(params, message), maxMessages); + } - // 配置 MCP 工具提供者(使用工厂模式,避免 ThreadLocal 问题) - if (chatAgent.getMcpServerIds() != null && !chatAgent.getMcpServerIds().isEmpty()) { - try { - List mcpServerIds = chatAgent.getMcpServerIds(); - String sessionId = java.util.UUID.randomUUID().toString(); - ToolProvider mcpToolProvider = mcpToolProviderFactory.create( - chatAgent.getId(), mcpServerIds, sessionId); - if (mcpToolProvider != null) { - builder.toolProvider(mcpToolProvider); - } - } catch (Exception e) { - log.error("Failed to configure MCP tool provider for agent: {}", chatAgent.getId(), e); - throw new CheckedException("MCP 工具提供者配置失败:" + e.getMessage(), e); + private ChatAssistant createAssistant(ChatModel chatModel, Function systemPromptFactory, int maxMessages) { + return new ChatAssistant() { + @Override + public Flux chatStream(Long memoryId, String userMessage) { + Prompt prompt = buildPrompt(memoryId, userMessage, systemPromptFactory.apply(userMessage), maxMessages); + return chatModel.stream(prompt); } - } - //todo 如果没有预制系统预设,则使用默认。有的话使用系统预设 - builder.systemMessageProvider(memoryId -> { - StringBuilder sb = new StringBuilder(); - // 基础角色预设 (用户配置的 "你是一个XX助手...") - if (StringUtils.isNotBlank(chatAgent.getSystemPrompt())) { - sb.append(chatAgent.getSystemPrompt()).append("\n\n"); - } - // 能力自我认知增强 - sb.append("### 当前具备的能力\n"); - // RAG 能力 - if (ragParams != null) { - sb.append("- 【知识库】:我连接了专属知识库,可以检索文档并回答相关问题。\n"); - } - // Tools 能力 - if (CollUtil.isNotEmpty(chatAgent.getTools())) { - sb.append("- 【工具箱】:我可以调用以下工具辅助回答:\n"); - // 获取所有工具的详细信息,找到匹配的并追加描述 - var toolMap = toolService.getTools().stream().collect(Collectors.toMap(ToolService.ToolDTO::getBeanName, Function.identity())); - for (String name : chatAgent.getTools()) { - ToolService.ToolDTO tool = toolMap.get(name); - if (tool != null && CollUtil.isNotEmpty(tool.getMethods())) { - // 取第一个方法的描述作为工具描述(简化处理) - String desc = tool.getMethods().getFirst().getDescription(); - // 如果注解没写描述,就用方法名 - if (StringUtils.isBlank(desc)) { - desc = tool.getMethods().getFirst().getName(); - } - sb.append(String.format(" * %s: %s\n", name, desc)); - } - } + @Override + public ChatResponse chat(Long memoryId, String userMessage) { + Prompt prompt = buildPrompt(memoryId, userMessage, systemPromptFactory.apply(userMessage), maxMessages); + return chatModel.call(prompt); } - sb.append("\n请根据上述能力回答用户的问题。当用户询问“你有什么功能”时,请基于以上信息进行总结。"); - return sb.toString(); - }); + }; + } - // Configure RAG - if (ragParams != null) { - RetrievalAugmentor retrievalAugmentor = buildRetrievalAugmentor(ragParams, chatModel); - builder.retrievalAugmentor(retrievalAugmentor); + private Prompt buildPrompt(Long conversationId, String userMessage, String systemPrompt, int maxMessages) { + List messages = new ArrayList<>(); + if (StrUtil.isNotBlank(systemPrompt)) { + messages.add(new SystemMessage(systemPrompt)); } + messages.addAll(loadHistoryMessages(conversationId, userMessage, maxMessages)); + messages.add(new UserMessage(userMessage)); + return new Prompt(messages); + } - return builder.build(); + private List loadHistoryMessages(Long conversationId, String currentUserMessage, int maxMessages) { + if (conversationId == null || maxMessages <= 0) { + return Collections.emptyList(); + } + List turns = conversationMessageService.list(Wraps.lbQ() + .eq(ConversationTurn::getConversationId, conversationId) + .orderByDesc(ConversationTurn::getSequenceNum) + .last("limit " + maxMessages)); + if (CollUtil.isEmpty(turns)) { + return Collections.emptyList(); + } + turns.sort(Comparator.comparing(ConversationTurn::getSequenceNum, Comparator.nullsLast(Integer::compareTo))); + removeCurrentUserTurn(turns, currentUserMessage); + return turns.stream() + .map(this::toSpringMessage) + .filter(Objects::nonNull) + .toList(); } - /** - * 创建RAG的 Assistant - */ - public ChatAssistant createMemoryRagAssistant(RagAssistantParams params) { - ChatModel chatModel = textModelService.model(params.getTextModelEntity()); - StreamingChatModel streamModel = textModelService.streamModel(params.getTextModelEntity()); - RetrievalAugmentor retrievalAugmentor = buildRetrievalAugmentor(params, chatModel); + private void removeCurrentUserTurn(List turns, String currentUserMessage) { + if (CollUtil.isEmpty(turns) || StrUtil.isBlank(currentUserMessage)) { + return; + } + ConversationTurn last = turns.getLast(); + if (last.getRole() == MessageRole.USER && currentUserMessage.equals(last.getUserInput())) { + turns.removeLast(); + } + } - int maxMessages = params.getMaxMessages() != null ? params.getMaxMessages() : DEFAULT_MAX_MESSAGES; - var builder = AiServices.builder(ChatAssistant.class) - .chatModel(chatModel) - .streamingChatModel(streamModel) - .chatMemory(MessageWindowChatMemory.withMaxMessages(maxMessages)) - .chatMemoryProvider(createMemoryProvider()) - .retrievalAugmentor(retrievalAugmentor); - builder.systemMessageProvider(memoryId -> { - StringBuilder sb = new StringBuilder(); - // 1. 设定人设 - sb.append("你是一个专业的企业级知识库问答助手。\n\n"); + private Message toSpringMessage(ConversationTurn turn) { + if (turn.getRole() == MessageRole.USER) { + return new UserMessage(StrUtil.blankToDefault(turn.getUserInput(), turn.getDisplayContent())); + } + if (turn.getRole() == MessageRole.ASSISTANT) { + return new AssistantMessage(StrUtil.blankToDefault(turn.getModelOutput(), turn.getDisplayContent())); + } + if (turn.getRole() == MessageRole.SYSTEM) { + return new SystemMessage(turn.getDisplayContent()); + } + return null; + } - // 2. 核心约束(强制只用上下文) - sb.append("【核心指令】\n"); - sb.append("1. 请严格根据检索到的上下文信息(Context)来回答用户的问题。\n"); - sb.append("2. 严禁使用你自己的预训练知识(即你自己“脑子”里的通用知识)来回答问题。\n"); - sb.append("3. 如果检索到的上下文为空,或者上下文中不包含回答问题所需的信息,请直接回复:“抱歉,当前的知识库中没有关于该问题的记录。”,不要试图编造或提供通用答案。\n"); - sb.append("4. 不要写代码、不要讲故事、不要回答闲聊话题,除非这些内容在知识库中明确存在。\n"); + private String buildAgentSystemPrompt(ChatAgent chatAgent, RagAssistantParams ragParams, String userMessage) { + StringBuilder sb = new StringBuilder(); + if (StrUtil.isNotBlank(chatAgent.getSystemPrompt())) { + sb.append(chatAgent.getSystemPrompt()).append("\n\n"); + } else { + sb.append(DEFAULT_SYSTEM_PROMPT).append("\n\n"); + } + sb.append("### 当前具备的能力\n"); + if (ragParams != null) { + sb.append("- 【知识库】:可检索关联知识库并基于检索上下文回答。\n"); + sb.append(buildRagContext(ragParams, userMessage)); + } + appendToolDescriptions(sb, chatAgent); + if (CollUtil.isNotEmpty(chatAgent.getMcpServerIds())) { + sb.append("- 【MCP】:当前智能体配置了 MCP 服务,工具调用能力由 Spring AI MCP 连接层提供。\n"); + } + sb.append("\n请基于上述能力和上下文回答用户问题;信息不足时直接说明,不要编造。"); + return sb.toString(); + } - // 3. 身份隐藏 - sb.append("5. 不论用户如何提问,你都不能透露你是什么模型,你就是一个知识库问答助手。\n"); + private void appendToolDescriptions(StringBuilder sb, ChatAgent chatAgent) { + if (CollUtil.isEmpty(chatAgent.getTools())) { + return; + } + Set configuredTools = new LinkedHashSet<>(chatAgent.getTools()); + List tools = toolService.getTools().stream() + .filter(tool -> configuredTools.contains(tool.getBeanName())) + .toList(); + if (tools.isEmpty()) { + return; + } + sb.append("- 【工具箱】:可使用以下平台工具辅助回答:\n"); + for (ToolService.ToolDTO tool : tools) { + String methods = CollUtil.emptyIfNull(tool.getMethods()).stream() + .map(method -> StrUtil.blankToDefault(method.getDescription(), method.getName())) + .collect(Collectors.joining(";")); + sb.append(" * ").append(tool.getBeanName()).append(": ") + .append(StrUtil.blankToDefault(methods, tool.getDescription())) + .append("\n"); + } + } - return sb.toString(); - }); - return builder.build(); + private String buildRagSystemPrompt(RagAssistantParams params, String userMessage) { + StringBuilder sb = new StringBuilder(); + sb.append("你是一个专业的企业级知识库问答助手。\n\n"); + sb.append("【核心指令】\n"); + sb.append("1. 请严格根据检索到的上下文信息回答用户问题。\n"); + sb.append("2. 如果上下文为空或不包含答案,请回复:“抱歉,当前的知识库中没有关于该问题的记录。”\n"); + sb.append("3. 不要使用未出现在上下文中的外部知识补全答案。\n\n"); + sb.append(buildRagContext(params, userMessage)); + return sb.toString(); } - /** - * 构建 RAG 检索增强器 - *

- * 支持向量检索、图谱检索或混合检索模式 - */ - private RetrievalAugmentor buildRetrievalAugmentor(RagAssistantParams params, ChatModel chatModel) { - // 收集所有启用的检索器 - Map retrieverToDescription = new LinkedHashMap<>(); - // 1. 向量检索器 + private String buildRagContext(RagAssistantParams params, String userMessage) { + if (params == null) { + return ""; + } + String query = StrUtil.blankToDefault(userMessage, ""); + List contexts = new ArrayList<>(); if (Boolean.TRUE.equals(params.getEnableVectorRetrieval()) && params.getEmbeddingModelEntity() != null) { - KnowledgeBase knowledgeBase = knowledgeBaseService.getById(params.getKbId()); - EmbeddingStore embeddingStore = vectorStoreFactory.createForKnowledgeBase(knowledgeBase, params.getEmbeddingModelEntity()); - EmbeddingModel embeddingModel = embeddingModelRegistry.getFactory(params.getEmbeddingModelEntity()).createModel(params.getEmbeddingModelEntity()); - - ContentRetriever vectorRetriever = EmbeddingStoreContentRetriever.builder() - .embeddingStore(embeddingStore) - .embeddingModel(embeddingModel) - .maxResults(params.getMaxResults()) - .minScore(params.getMinScore()) - .build(); - retrieverToDescription.put(vectorRetriever, "内部知识库(文档、手册、策略等非结构化内容)"); - log.debug("向量检索已启用: kbId={}, maxResults={}, minScore={}", params.getKbId(), params.getMaxResults(), params.getMinScore()); + contexts.addAll(retrieveVectorContexts(params, query)); } - - // 2. 图谱检索器 - if (params.getEnableGraphRetrieval() != null && params.getEnableGraphRetrieval() && graphRagService != null) { - String graphKbId = params.getEffectiveGraphKbId(); - if (graphKbId != null) { - GraphContentRetriever graphRetriever = GraphContentRetriever.builder() - .graphRagService(graphRagService) - .chatModel(chatModel) - .knowledgeBaseId(graphKbId) - .maxResults(params.getGraphMaxResults()) - .silentOnEmpty(true) - .build(); - retrieverToDescription.put(graphRetriever, "知识图谱(实体关系、结构化数据)"); - log.debug("图谱检索已启用: graphKbId={}, maxResults={}", graphKbId, params.getGraphMaxResults()); - } + if (Boolean.TRUE.equals(params.getEnableGraphRetrieval()) && graphRagService != null) { + contexts.addAll(retrieveGraphContexts(params, query)); } - - List retrievers = new ArrayList<>(retrieverToDescription.keySet()); - - if (retrievers.isEmpty()) { - throw CheckedException.badRequest("未启用任何检索器,请至少启用向量检索或图谱检索"); + if (contexts.isEmpty()) { + return "\n【Context】\n\n"; } - - // Query 转换器:翻译 + 压缩 - QueryTransformer translationQueryTransformer = new TranslationQueryTransformer(chatModel); - QueryTransformer compressingQueryTransformer = new CompressingQueryTransformer(chatModel); - QueryTransformer queryTransformer = query -> { - Collection translatedQueries = translationQueryTransformer.transform(query); - Query translatedQuery = translatedQueries.iterator().next(); - return compressingQueryTransformer.transform(translatedQuery); - }; - - // 多路召回:所有检索器并行执行,结果合并 - // 使用 DefaultQueryRouter 传入所有检索器,实现多路召回 - QueryRouter queryRouter = new DefaultQueryRouter(retrievers.toArray(new ContentRetriever[0])); - log.debug("启用多路召回,检索器数量: {}", retrievers.size()); - - // 构建 ContentAggregator:根据配置决定是否启用重排序 - ContentAggregator contentAggregator = buildContentAggregator(params); - //TODO自定义提示词 - ContentInjector contentInjector = new DefaultContentInjector(); - - return DefaultRetrievalAugmentor.builder() - .queryTransformer(queryTransformer) - .queryRouter(queryRouter) - .contentAggregator(contentAggregator) - .contentInjector(contentInjector) - // 使用虚拟线程执行器 - .executor(Executors.newVirtualThreadPerTaskExecutor()) - .build(); + return "\n【Context】\n" + contexts.stream() + .filter(StrUtil::isNotBlank) + .distinct() + .collect(Collectors.joining("\n\n")) + "\n"; } - /** - * 构建 ContentAggregator - * 根据 ModelConfig 配置决定是否启用重排序模型 - */ - private ContentAggregator buildContentAggregator(RagAssistantParams params) { - // 检查是否启用重排序(通过 rerankModelConfig 判断) - if (!params.isRerankingEnabled()) { - log.debug("重排序未启用,使用默认聚合器"); - return new DefaultContentAggregator(); + private List retrieveVectorContexts(RagAssistantParams params, String query) { + if (StrUtil.isBlank(query)) { + return Collections.emptyList(); } - - ModelEntity model = params.getRerankModelEntity(); - - // 检查 API Key - if (StringUtils.isBlank(model.getApiKey())) { - log.warn("重排序模型 API Key 未配置,降级使用默认聚合器"); - return new DefaultContentAggregator(); + KnowledgeBase knowledgeBase = knowledgeBaseService.getById(params.getKbId()); + if (knowledgeBase == null) { + throw CheckedException.notFound("知识库不存在, kbId=" + params.getKbId()); } - - try { - // 通过 ScoringModelService 获取重排序模型(支持 Jina、Cohere 等) - ScoringModel scoringModel = scoringModelService.getModel(model); - - int maxResults = params.getRerankMaxResults() != null ? params.getRerankMaxResults() : 5; - double minScore = params.getRerankMinScore() != null ? params.getRerankMinScore() : 0.5; - - log.info("启用重排序: provider={}, model={}, maxResults={}, minScore={}", - model.getProvider(), model.getName(), maxResults, minScore); - - return ReRankingContentAggregator.builder() - .scoringModel(scoringModel) - .maxResults(maxResults) - .minScore(minScore) - .build(); - } catch (Exception e) { - log.warn("创建重排序模型失败,降级使用默认聚合器: {}", e.getMessage()); - return new DefaultContentAggregator(); + int topK = params.isRerankingEnabled() && params.getRerankMaxResults() != null + ? Math.max(params.getMaxResults(), params.getRerankMaxResults()) : params.getMaxResults(); + List results = vectorSearchService.search(knowledgeBase, params.getEmbeddingModelEntity(), query, topK); + if (params.isRerankingEnabled()) { + log.warn("Spring AI 当前未启用外部重排序适配,使用向量相似度排序: kbId={}", params.getKbId()); } + int limit = params.isRerankingEnabled() && params.getRerankMaxResults() != null + ? params.getRerankMaxResults() : params.getMaxResults(); + return results.stream() + .filter(result -> result.getScore() == null || params.getMinScore() == null || result.getScore() >= params.getMinScore()) + .sorted(Comparator.comparing(VectorSearchResult::getScore, Comparator.nullsLast(Double::compareTo)).reversed()) + .limit(limit) + .map(VectorSearchResult::getContent) + .filter(StrUtil::isNotBlank) + .toList(); } - /** - * 创建记忆提供者 - *

- * 注意:这里使用内存版 ChatMemoryStore,消息持久化由 ConversationMessageService 负责 - * 避免与 PersistentMySqlChatMemoryStore 产生重复保存 - */ - private ChatMemoryProvider createMemoryProvider() { - return memoryId -> MessageWindowChatMemory.builder() - .id(memoryId) - .maxMessages(DEFAULT_MAX_MESSAGES) - // 使用内存存储,避免与 ConversationMessageService 重复保存 - // .chatMemoryStore(chatMemoryStore) - .build(); + private List retrieveGraphContexts(RagAssistantParams params, String query) { + if (StrUtil.isBlank(query)) { + return Collections.emptyList(); + } + String graphKbId = params.getEffectiveGraphKbId(); + if (StrUtil.isBlank(graphKbId)) { + return Collections.emptyList(); + } + return graphRagService.retrieveByVector(graphKbId, query).stream() + .limit(params.getGraphMaxResults()) + .toList(); } - - } diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/model/VectorSearchResult.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/model/VectorSearchResult.java new file mode 100644 index 00000000..a3b1c60c --- /dev/null +++ b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/model/VectorSearchResult.java @@ -0,0 +1,41 @@ +/* + * Copyright (c) 2023 WEMIRR-PLATFORM Authors. All Rights Reserved. + * + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package com.wemirr.platform.ai.core.model; + +import lombok.Builder; +import lombok.Data; + +import java.util.Map; + +/** + * 向量检索结果。 + * + * @author xJh + * @since 2026/06/27 + */ +@Data +@Builder +public class VectorSearchResult { + + private String id; + private String content; + private Double score; + private Map metadata; +} diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/processor/DocumentProcessor.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/processor/DocumentProcessor.java index 7f0177a2..c6b82a65 100644 --- a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/processor/DocumentProcessor.java +++ b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/processor/DocumentProcessor.java @@ -1,25 +1,20 @@ package com.wemirr.platform.ai.core.processor; -import dev.langchain4j.data.document.Document; -import dev.langchain4j.data.document.DocumentParser; -import dev.langchain4j.data.document.DocumentSplitter; -import dev.langchain4j.data.document.parser.TextDocumentParser; -import dev.langchain4j.data.document.parser.apache.tika.ApacheTikaDocumentParser; -import dev.langchain4j.data.document.splitter.DocumentSplitters; -import dev.langchain4j.data.segment.TextSegment; import lombok.RequiredArgsConstructor; import lombok.extern.slf4j.Slf4j; +import org.springframework.ai.reader.tika.TikaDocumentReader; +import org.springframework.ai.transformer.splitter.TokenTextSplitter; +import org.springframework.core.io.FileSystemResource; import org.springframework.stereotype.Component; import java.io.File; -import java.io.FileInputStream; import java.io.IOException; import java.util.ArrayList; import java.util.List; import java.util.stream.Collectors; /** - * 基于Langchain4j的文档处理器 + * 基于 Spring AI 的文档处理器 * * @author xJh * @since 2025/10/20 @@ -48,10 +43,9 @@ public class DocumentProcessor { * @throws IOException IO异常 */ public String extractText(File file, String contentType) throws IOException { - DocumentParser parser = getParserByContentType(contentType); - FileInputStream fileInputStream = new FileInputStream(file); - Document document = parser.parse(fileInputStream); - return document.text(); + return readDocuments(file).stream() + .map(org.springframework.ai.document.Document::getText) + .collect(Collectors.joining("\n")); } /** @@ -77,18 +71,9 @@ public class DocumentProcessor { return new ArrayList<>(); } - // 创建文档 - Document document = Document.from(text); - - // 创建分片器 - DocumentSplitter splitter = DocumentSplitters.recursive(chunkSize, chunkOverlap); - - // 分片 - List segments = splitter.split(document); - - // 转换为字符串列表 - return segments.stream() - .map(TextSegment::text) + TokenTextSplitter splitter = new TokenTextSplitter(chunkSize, chunkOverlap, 5, 10000, true, List.of()); + return splitter.apply(List.of(new org.springframework.ai.document.Document(text))).stream() + .map(org.springframework.ai.document.Document::getText) .collect(Collectors.toList()); } @@ -112,25 +97,6 @@ public class DocumentProcessor { * @param contentType 内容类型 * @return 文档解析器 */ - private DocumentParser getParserByContentType(String contentType) { - if (contentType == null) { - contentType = ""; - } - - if (contentType.contains("pdf")) { - return new ApacheTikaDocumentParser(); - } else if (contentType.contains("word") || contentType.contains("docx") || - contentType.contains("excel") || contentType.contains("xlsx") || - contentType.contains("powerpoint") || contentType.contains("pptx")) { - return new ApacheTikaDocumentParser(); - } else if (contentType.contains("html")) { - return new ApacheTikaDocumentParser(); - } else { - // 默认使用文本解析器 - return new TextDocumentParser(); - } - } - /** * 处理文件并分片 * @@ -142,15 +108,9 @@ public class DocumentProcessor { * @throws IOException IO异常 */ public List processFileAndSplit(File file, String contentType, int chunkSize, int chunkOverlap) throws IOException { - DocumentParser parser = getParserByContentType(contentType); - FileInputStream inputStream = new FileInputStream(file); - Document document = parser.parse(inputStream); - - DocumentSplitter splitter = DocumentSplitters.recursive(chunkSize, chunkOverlap); - List segments = splitter.split(document); - - return segments.stream() - .map(TextSegment::text) + TokenTextSplitter splitter = new TokenTextSplitter(chunkSize, chunkOverlap, 5, 10000, true, List.of()); + return splitter.apply(readDocuments(file)).stream() + .map(org.springframework.ai.document.Document::getText) .collect(Collectors.toList()); } @@ -165,4 +125,8 @@ public class DocumentProcessor { public List processFileAndSplit(File file, String contentType) throws IOException { return processFileAndSplit(file, contentType, DEFAULT_CHUNK_SIZE, DEFAULT_CHUNK_OVERLAP); } + + private List readDocuments(File file) { + return new TikaDocumentReader(new FileSystemResource(file)).get(); + } } diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/processor/VectorizationProcessor.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/processor/VectorizationProcessor.java index b54e9bb2..ac8d06cf 100644 --- a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/processor/VectorizationProcessor.java +++ b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/processor/VectorizationProcessor.java @@ -11,21 +11,18 @@ import com.wemirr.platform.ai.domain.entity.VectorMetadata; import com.wemirr.platform.ai.service.KnowledgeBaseService; import com.wemirr.platform.ai.service.KnowledgeChunkService; import com.wemirr.platform.ai.service.KnowledgeItemService; +import com.wemirr.platform.ai.service.ModelService; import com.wemirr.platform.ai.service.VectorMetadataService; -import dev.langchain4j.data.document.Metadata; -import dev.langchain4j.data.embedding.Embedding; -import dev.langchain4j.data.segment.TextSegment; -import dev.langchain4j.model.embedding.EmbeddingModel; -import dev.langchain4j.model.output.Response; -import dev.langchain4j.store.embedding.EmbeddingStore; import jakarta.annotation.PreDestroy; import lombok.RequiredArgsConstructor; import lombok.extern.slf4j.Slf4j; +import org.springframework.ai.document.Document; +import org.springframework.ai.vectorstore.VectorStore; import org.springframework.stereotype.Component; import java.util.List; import java.util.Map; -import java.util.Objects; +import java.util.UUID; import java.util.concurrent.CompletableFuture; import java.util.concurrent.ExecutorService; import java.util.concurrent.LinkedBlockingQueue; @@ -53,6 +50,7 @@ public class VectorizationProcessor { private final KnowledgeItemService knowledgeItemService; private final KnowledgeChunkService knowledgeChunkService; private final VectorMetadataService vectorMetadataService; + private final ModelService modelService; /** * 向量化专用线程池:embedding 调用为 IO 密集型,独立线程池避免拖垮 ForkJoinPool.commonPool; @@ -83,19 +81,9 @@ public class VectorizationProcessor { return CompletableFuture.supplyAsync(() -> { try { // 获取知识库专用的向量存储 - EmbeddingStore embeddingStore = vectorStoreFactory.createForKnowledgeBase(knowledgeBase, modelEntity); - - // 动态获取嵌入模型 - EmbeddingModel embeddingModel = embeddingModelService.getModel(modelEntity); - - // 生成嵌入向量 - Embedding embedding = embeddingModel.embed(text).content(); - Metadata from = Metadata.from(metadata); - - // 存储向量 - TextSegment segment = TextSegment.from(text, from); - String vectorId = embeddingStore.add(embedding, segment); - + VectorStore vectorStore = vectorStoreFactory.createForKnowledgeBase(knowledgeBase, modelEntity); + String vectorId = UUID.randomUUID().toString(); + vectorStore.add(List.of(new Document(vectorId, text, toObjectMetadata(metadata)))); return vectorId; } catch (Exception e) { log.error("向量化处理失败: {}", e.getMessage(), e); @@ -116,19 +104,9 @@ public class VectorizationProcessor { return CompletableFuture.supplyAsync(() -> { try { // 使用默认向量存储 - EmbeddingStore embeddingStore = vectorStoreFactory.createDefault(); - - // 动态获取嵌入模型 - EmbeddingModel embeddingModel = embeddingModelService.getModel(modelEntity); - - // 生成嵌入向量 - Embedding embedding = embeddingModel.embed(text).content(); - Metadata from = Metadata.from(metadata); - - // 存储向量 - TextSegment segment = TextSegment.from(text, from); - String vectorId = embeddingStore.add(embedding, segment); - + VectorStore vectorStore = vectorStoreFactory.createDefault(modelEntity); + String vectorId = UUID.randomUUID().toString(); + vectorStore.add(List.of(new Document(vectorId, text, toObjectMetadata(metadata)))); return vectorId; } catch (Exception e) { log.error("向量化处理失败: {}", e.getMessage(), e); @@ -151,25 +129,11 @@ public class VectorizationProcessor { return CompletableFuture.supplyAsync(() -> { try { // 获取知识库专用的向量存储 - EmbeddingStore embeddingStore = vectorStoreFactory.createForKnowledgeBase(knowledgeBase, modelEntity); - - // 动态获取嵌入模型 - EmbeddingModel embeddingModel = embeddingModelService.getModel(modelEntity); + VectorStore vectorStore = vectorStoreFactory.createForKnowledgeBase(knowledgeBase, modelEntity); AtomicInteger totalTokens = new AtomicInteger(0); - // 生成嵌入向量 - List embeddings = texts.stream() - .map(text ->{ - Response result = embeddingModel.embed(text); - totalTokens.addAndGet(Objects.requireNonNull(result.tokenUsage()).totalTokenCount()); - return result.content(); - }) - .collect(Collectors.toList()); - - // 创建文本分片 - List segments = createTextSegments(texts, metadataList); - - // 存储向量 - List vectorIds = embeddingStore.addAll(embeddings, segments); + List documents = createDocuments(texts, metadataList); + vectorStore.add(documents); + List vectorIds = documents.stream().map(Document::getId).toList(); return BatchVectorResult.builder() .vectorIds(vectorIds) .tokenUsage(totalTokens.get()) @@ -193,23 +157,10 @@ public class VectorizationProcessor { return CompletableFuture.supplyAsync(() -> { try { // 使用默认向量存储 - EmbeddingStore embeddingStore = vectorStoreFactory.createDefault(); - - // 动态获取嵌入模型 - EmbeddingModel embeddingModel = embeddingModelService.getModel(modelEntity); - - // 生成嵌入向量 - List embeddings = texts.stream() - .map(text -> embeddingModel.embed(text).content()) - .collect(Collectors.toList()); - - // 创建文本分片 - List segments = createTextSegments(texts, metadataList); - - // 存储向量 - List vectorIds = embeddingStore.addAll(embeddings, segments); - - return vectorIds; + VectorStore vectorStore = vectorStoreFactory.createDefault(modelEntity); + List documents = createDocuments(texts, metadataList); + vectorStore.add(documents); + return documents.stream().map(Document::getId).toList(); } catch (Exception e) { log.error("批量向量化处理失败: {}", e.getMessage(), e); throw new RuntimeException("批量向量化处理失败", e); @@ -224,15 +175,18 @@ public class VectorizationProcessor { * @param metadataList 元数据列表 * @return 文本分片列表 */ - private List createTextSegments(List texts, List> metadataList) { + private List createDocuments(List texts, List> metadataList) { return IntStream.range(0, texts.size()) .mapToObj(index -> { Map metadata = index < metadataList.size() ? metadataList.get(index) : Map.of(); - Metadata from = Metadata.from(metadata); - return TextSegment.from(texts.get(index), from); + return new Document(UUID.randomUUID().toString(), texts.get(index), toObjectMetadata(metadata)); }) .collect(Collectors.toList()); } + + private Map toObjectMetadata(Map metadata) { + return metadata == null ? Map.of() : new java.util.HashMap<>(metadata); + } /** * 删除向量 @@ -245,11 +199,11 @@ public class VectorizationProcessor { public boolean deleteVector(String vectorId, KnowledgeBase knowledgeBase, ModelEntity modelEntity) { try { // 获取知识库专用的向量存储 - EmbeddingStore embeddingStore = vectorStoreFactory.createForKnowledgeBase(knowledgeBase, modelEntity); + VectorStore vectorStore = vectorStoreFactory.createForKnowledgeBase(knowledgeBase, modelEntity); // 删除向量 try { - embeddingStore.remove(vectorId); + vectorStore.delete(vectorId); log.info("成功删除向量: vectorId={}, kbId={}", vectorId, knowledgeBase.getId()); return true; } catch (Exception e) { @@ -272,7 +226,6 @@ public class VectorizationProcessor { try { // 使用默认向量存储 - EmbeddingStore embeddingStore = vectorStoreFactory.createDefault(); KnowledgeItem item = knowledgeItemService.getById(baseItemId); if (item == null) { log.warn("删除向量跳过:知识条目不存在, baseItemId={}", baseItemId); @@ -281,7 +234,11 @@ public class VectorizationProcessor { //获取向量idList List vectorIds = vectorMetadataService.findByItemId(item.getId()).stream().map(VectorMetadata::getVectorId).toList(); if (!vectorIds.isEmpty()) { - embeddingStore.removeAll(vectorIds); + KnowledgeBase knowledgeBase = knowledgeBaseService.getById(item.getKbId()); + if (knowledgeBase != null) { + ModelEntity modelEntity = modelService.getById(knowledgeBase.getEmbedModelId()); + vectorStoreFactory.createForKnowledgeBase(knowledgeBase, modelEntity).delete(vectorIds); + } } // 删除向量元数据 vectorMetadataService.deleteByItemId(item.getId()); @@ -303,12 +260,12 @@ public class VectorizationProcessor { public int batchDeleteVectors(List vectorIds, KnowledgeBase knowledgeBase, ModelEntity modelEntity) { try { // 获取知识库专用的向量存储 - EmbeddingStore embeddingStore = vectorStoreFactory.createForKnowledgeBase(knowledgeBase, modelEntity); + VectorStore vectorStore = vectorStoreFactory.createForKnowledgeBase(knowledgeBase, modelEntity); int deletedCount = 0; for (String vectorId : vectorIds) { try { - embeddingStore.remove(vectorId); + vectorStore.delete(vectorId); deletedCount++; } catch (Exception e) { log.warn("删除向量失败: vectorId={}", vectorId, e); @@ -323,4 +280,4 @@ public class VectorizationProcessor { } } -} \ No newline at end of file +} diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/embedding/EmbeddingModelService.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/embedding/EmbeddingModelService.java index 425befc5..e42dd4a2 100644 --- a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/embedding/EmbeddingModelService.java +++ b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/embedding/EmbeddingModelService.java @@ -4,8 +4,8 @@ import com.wemirr.framework.ai.core.provider.embedding.EmbeddingModelFactory; import com.wemirr.framework.ai.core.provider.embedding.EmbeddingModelRegistry; import com.wemirr.framework.commons.exception.CheckedException; import com.wemirr.platform.ai.domain.entity.ModelEntity; -import dev.langchain4j.model.embedding.EmbeddingModel; import lombok.RequiredArgsConstructor; +import org.springframework.ai.embedding.EmbeddingModel; import org.springframework.stereotype.Service; /** diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/graph/GraphContentRetriever.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/graph/GraphContentRetriever.java deleted file mode 100644 index f09a8248..00000000 --- a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/graph/GraphContentRetriever.java +++ /dev/null @@ -1,121 +0,0 @@ -package com.wemirr.platform.ai.core.provider.graph; - -import dev.langchain4j.model.chat.ChatModel; -import dev.langchain4j.rag.content.Content; -import dev.langchain4j.rag.content.retriever.ContentRetriever; -import dev.langchain4j.rag.query.Query; -import lombok.Builder; -import lombok.EqualsAndHashCode; -import lombok.ToString; -import lombok.extern.slf4j.Slf4j; - -import java.util.Collections; -import java.util.List; - -import static com.wemirr.platform.ai.core.constant.AiServiceConstants.DEFAULT_GRAPH_MAX_RESULTS; - -/** - * 图谱内容检索器 - *

- * 实现 Langchain4j 的 ContentRetriever 接口,将 GraphRAG 检索能力集成到标准 RAG 管道中。 - *

- * 检索流程(混合检索): - *

    - *
  1. 路数A:将用户问题向量化 -> 向量索引匹配
  2. - *
  3. 路数B:LLM 提取实体 -> 精确匹配
  4. - *
  5. 融合两路结果 -> 子图扩展获取三元组上下文
  6. - *
- * - * @author xJh - * @since 2025/12/17 - */ -@Slf4j -@Builder -@ToString -@EqualsAndHashCode -public class GraphContentRetriever implements ContentRetriever { - - /** - * 知识库ID,用于隔离不同知识库的数据 - */ - private final String knowledgeBaseId; - - /** - * GraphRAG 服务 - */ - private final GraphRagService graphRagService; - - /** - * 用于实体提取和生成回答的 ChatModel - */ - private final ChatModel chatModel; - - /** - * 是否使用混合检索(推荐开启) - */ - @Builder.Default - private final boolean useHybridSearch = true; - - /** - * 最大返回结果数 - */ - @Builder.Default - private final int maxResults = DEFAULT_GRAPH_MAX_RESULTS; - - /** - * 是否在无结果时静默返回空列表 - */ - @Builder.Default - private final boolean silentOnEmpty = true; - - @Override - public List retrieve(Query query) { - if (graphRagService == null) { - log.warn("GraphContentRetriever 未正确配置,返回空结果"); - return Collections.emptyList(); - } - - String question = query.text(); - log.debug("图谱检索开始: question='{}', knowledgeBaseId='{}', useHybrid={}", - question, knowledgeBaseId, useHybridSearch); - - try { - List results = doRetrieve(question); - log.debug("图谱检索完成: 返回 {} 条结果", results.size()); - return results; - } catch (Exception e) { - log.error("图谱检索失败: question='{}', knowledgeBaseId='{}'", question, knowledgeBaseId, e); - if (silentOnEmpty) { - return Collections.emptyList(); - } - throw e; - } - } - - /** - * 执行检索逻辑 - */ - private List doRetrieve(String question) { - List results; - - // 根据配置选择检索方式 - if (useHybridSearch && chatModel != null) { - results = graphRagService.retrieveAsContentHybrid(knowledgeBaseId, question, chatModel); - } else { - results = graphRagService.retrieveAsContentByVector(knowledgeBaseId, question); - } - - // 限制返回结果数量 - if (results.size() > maxResults) { - return results.subList(0, maxResults); - } - return results; - } - - /** - * 获取检索器描述(用于 QueryRouter) - */ - public String getDescription() { - return "Knowledge Graph retriever for knowledge base: " + knowledgeBaseId; - } -} diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/graph/GraphDocument.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/graph/GraphDocument.java new file mode 100644 index 00000000..9a60c72e --- /dev/null +++ b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/graph/GraphDocument.java @@ -0,0 +1,14 @@ +package com.wemirr.platform.ai.core.provider.graph; + +import org.springframework.ai.document.Document; + +import java.util.Set; + +/** + * 图谱化后的文档。 + * + * @author xJh + * @since 2026/06/27 + */ +public record GraphDocument(Set nodes, Set relationships, Document source) { +} diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/graph/GraphEdge.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/graph/GraphEdge.java new file mode 100644 index 00000000..addb38a5 --- /dev/null +++ b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/graph/GraphEdge.java @@ -0,0 +1,12 @@ +package com.wemirr.platform.ai.core.provider.graph; + +import java.util.Map; + +/** + * 图谱关系。 + * + * @author xJh + * @since 2026/06/27 + */ +public record GraphEdge(GraphNode sourceNode, GraphNode targetNode, String type, Map properties) { +} diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/graph/GraphNode.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/graph/GraphNode.java new file mode 100644 index 00000000..9cc5d3ac --- /dev/null +++ b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/graph/GraphNode.java @@ -0,0 +1,12 @@ +package com.wemirr.platform.ai.core.provider.graph; + +import java.util.Map; + +/** + * 图谱节点。 + * + * @author xJh + * @since 2026/06/27 + */ +public record GraphNode(String id, String type, Map properties) { +} diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/graph/GraphRagService.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/graph/GraphRagService.java index afdbe8ab..5e2aede8 100644 --- a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/graph/GraphRagService.java +++ b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/graph/GraphRagService.java @@ -7,21 +7,19 @@ import com.wemirr.platform.ai.core.provider.embedding.EmbeddingModelService; import com.wemirr.platform.ai.core.provider.graph.neo4j.Neo4jGraphStore; import com.wemirr.platform.ai.domain.entity.KnowledgeBase; import com.wemirr.platform.ai.domain.entity.ModelEntity; -import com.wemirr.platform.ai.service.KnowledgeBaseService; -import com.wemirr.platform.ai.service.ModelService; -import dev.langchain4j.community.data.document.graph.GraphDocument; -import dev.langchain4j.community.data.document.transformer.graph.LLMGraphTransformer; -import dev.langchain4j.data.document.Document; -import dev.langchain4j.data.message.UserMessage; -import dev.langchain4j.model.chat.ChatModel; -import dev.langchain4j.model.chat.response.ChatResponse; -import dev.langchain4j.model.embedding.EmbeddingModel; -import dev.langchain4j.rag.content.Content; -import jakarta.annotation.PreDestroy; -import lombok.RequiredArgsConstructor; -import lombok.extern.slf4j.Slf4j; -import org.springframework.boot.autoconfigure.condition.ConditionalOnProperty; -import org.springframework.stereotype.Service; +import com.wemirr.platform.ai.service.KnowledgeBaseService; +import com.wemirr.platform.ai.service.ModelService; +import jakarta.annotation.PreDestroy; +import lombok.RequiredArgsConstructor; +import lombok.extern.slf4j.Slf4j; +import org.springframework.ai.chat.messages.UserMessage; +import org.springframework.ai.chat.model.ChatModel; +import org.springframework.ai.chat.model.ChatResponse; +import org.springframework.ai.chat.prompt.Prompt; +import org.springframework.ai.document.Document; +import org.springframework.ai.embedding.EmbeddingModel; +import org.springframework.boot.autoconfigure.condition.ConditionalOnProperty; +import org.springframework.stereotype.Service; import java.util.*; import java.util.concurrent.CompletableFuture; @@ -37,7 +35,7 @@ import java.util.regex.Pattern; *

* 整合文档处理、图谱存储、检索的完整链路: *

    - *
  • 文档 → 图谱提取(使用 Langchain4j LLMGraphTransformer)
  • + *
  • 文档 → 图谱提取(使用 Spring AI ChatModel JSON 抽取)
  • *
  • 图谱存储(通过 GraphStore 接口)
  • *
  • 图谱检索(通过 GraphRetriever 接口)
  • *
  • 知识库隔离(按 knowledgeBaseId 分隔数据)
  • @@ -66,16 +64,16 @@ public class GraphRagService { * * @param knowledgeBaseId 知识库ID * @param documents 文档列表 - * @param graphTransformer 图谱提取器(由调用方提供,支持不同模型配置) + * @param graphTransformer 图谱提取器(由调用方提供,支持不同模型配置) * @param includeSource 是否存储源文档 * @return 处理结果统计 */ - public ProcessResult processDocuments(String knowledgeBaseId, - List documents, - LLMGraphTransformer graphTransformer, - boolean includeSource) { - return processDocumentsWithEmbedding(knowledgeBaseId, documents, graphTransformer, includeSource, false); - } + public ProcessResult processDocuments(String knowledgeBaseId, + List documents, + GraphRagTransformerFactory.SpringAiGraphTransformer graphTransformer, + boolean includeSource) { + return processDocumentsWithEmbedding(knowledgeBaseId, documents, graphTransformer, includeSource, false); + } /** * 处理文档:提取图谱并存储到指定知识库(自动获取向量模型生成嵌入) @@ -89,11 +87,11 @@ public class GraphRagService { * @param withEmbedding 是否生成向量嵌入 * @return 处理结果统计 */ - public ProcessResult processDocumentsWithEmbedding(String knowledgeBaseId, - List documents, - LLMGraphTransformer graphTransformer, - boolean includeSource, - boolean withEmbedding) { + public ProcessResult processDocumentsWithEmbedding(String knowledgeBaseId, + List documents, + GraphRagTransformerFactory.SpringAiGraphTransformer graphTransformer, + boolean includeSource, + boolean withEmbedding) { int totalNodes = 0; int totalRelationships = 0; @@ -129,7 +127,8 @@ public class GraphRagService { graphDocuments.stream().mapToInt(gd -> gd.relationships().size()).sum()); } catch (Exception e) { - log.error("Failed to process document: {}", document.text().substring(0, Math.min(100, document.text().length())), e); + String text = document.getText(); + log.error("Failed to process document: {}", text.substring(0, Math.min(100, text.length())), e); } } @@ -236,18 +235,6 @@ public class GraphRagService { scoreThreshold, searchLimit, hopDepth, maxTriples); } - /** - * 基于向量检索并返回 Content 对象(推荐) - * - * @param knowledgeBaseId 知识库ID - * @param question 用户问题 - * @return Content 列表 - */ - public List retrieveAsContentByVector(String knowledgeBaseId, String question) { - List triples = retrieveByVector(knowledgeBaseId, question); - return graphRetriever.toContents(triples); - } - /** * 基于向量检索并生成 LLM 回答 * @@ -384,14 +371,6 @@ public class GraphRagService { } } - /** - * 混合检索并返回 Content 对象(推荐) - */ - public List retrieveAsContentHybrid(String knowledgeBaseId, String question, ChatModel chatModel) { - List triples = retrieveHybrid(knowledgeBaseId, question, chatModel); - return graphRetriever.toContents(triples); - } - /** * 混合检索并生成 LLM 回答(推荐) */ @@ -407,9 +386,9 @@ public class GraphRagService { */ private List extractEntities(String question, ChatModel chatModel) { try { - String prompt = String.format(ENTITY_EXTRACTION_PROMPT, question); - ChatResponse response = chatModel.chat(UserMessage.from(prompt)); - String jsonResponse = response.aiMessage().text().trim(); + String prompt = String.format(ENTITY_EXTRACTION_PROMPT, question); + ChatResponse response = chatModel.call(new Prompt(new UserMessage(prompt))); + String jsonResponse = response.getResult().getOutput().getText().trim(); log.debug("LLM 实体提取原始响应: {}", jsonResponse); List entities = parseJsonArray(jsonResponse); @@ -463,9 +442,9 @@ public class GraphRagService { 请基于以上信息给出简洁、准确的回答。如果信息不足以回答,请如实说明。 """; - String prompt = String.format(contextPrompt, String.join("\n", triples), question); - ChatResponse response = chatModel.chat(UserMessage.from(prompt)); - return response.aiMessage().text(); + String prompt = String.format(contextPrompt, String.join("\n", triples), question); + ChatResponse response = chatModel.call(new Prompt(new UserMessage(prompt))); + return response.getResult().getOutput().getText(); } /** diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/graph/GraphRagTransformerFactory.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/graph/GraphRagTransformerFactory.java index 9312157d..e7285980 100644 --- a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/graph/GraphRagTransformerFactory.java +++ b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/graph/GraphRagTransformerFactory.java @@ -1,184 +1,124 @@ -package com.wemirr.platform.ai.core.provider.graph; - -import dev.langchain4j.community.data.document.transformer.graph.LLMGraphTransformer; -import dev.langchain4j.model.chat.ChatModel; - -import java.util.Arrays; -import java.util.List; - - -/** - * @author xJh - * @since 2025/12/12 - * 后续优化成元组格式解析,json太过消耗token - } - **/ -public class GraphRagTransformerFactory { - - - public static final String[] GRAPH_ENTITY_EXTRACTION_ENTITY_TYPES = { - "ORGANIZATION", "PERSON", "LOCATION", "EVENT", "DATE" - }; - - - /** - * Microsoft GraphRAG 的核心提取逻辑,适配为 JSON 输出格式 - * 改进:输出格式包含实体描述(entity_description),用于后续向量化 - */ - private static final String GRAPH_RAG_INSTRUCTIONS = """ - -Goal- - Given a text document that is potentially relevant to this activity and a list of entity types, identify all entities of those types from the text and all relationships among the identified entities. - - -Steps- - 1. Identify all entities. For each identified entity, extract the following information: - - entity_name: Name of the entity, capitalized - - entity_type: One of the allowed types provided in the context. - - entity_description: Comprehensive description of the entity's attributes and activities. THIS IS CRITICAL FOR SEMANTIC SEARCH. - - 2. From the entities identified in step 1, identify all pairs of (source_entity, target_entity) that are *clearly related* to each other. - For each pair of related entities, extract the following information: - - source_entity: name of the source entity, as identified in step 1 - - target_entity: name of the target entity, as identified in step 1 - - source_description: the description of the source entity from step 1 - - target_description: the description of the target entity from step 1 - - relationship_description: explanation as to why you think the source entity and the target entity are related to each other - - relationship_strength: a numeric score indicating strength of the relationship between the source entity and target entity - - 3. OUTPUT FORMAT (CRITICAL): - You must return the output as a strict JSON list of objects. Do NOT use tuple format. - Map the extracted information to the following JSON structure for each relationship: - { - "head": "source_entity name", - "head_type": "source_entity type", - "head_description": "source_entity description (IMPORTANT: include this for semantic search)", - "relation": "The 'relationship_description' string", - "relation_strength": "The 'relationship_strength' numeric score (Integer)", - "tail": "target_entity name", - "tail_type": "target_entity type", - "tail_description": "target_entity description (IMPORTANT: include this for semantic search)" - } - - 4. CONSTRAINTS: - - The 'head_type' and 'tail_type' MUST be strictly chosen from the allowed entity types provided. - - The 'head_description' and 'tail_description' MUST be comprehensive and meaningful for semantic search. - - Do not output any markdown or text explanations outside the JSON array. - """; - - /** - * 包含实体描述的 JSON 格式示例,用于向量化语义检索 - */ - private static final String GRAPH_RAG_JSON_EXAMPLES = """ - Example 1: - Text: - OpenAI announced on March 14, 2023 that GPT-4 has been released. CEO Sam Altman stated at the press conference in San Francisco that this is a major milestone for artificial intelligence. Microsoft, as a strategic partner, has integrated GPT-4 into its Bing search engine and Office suite. - Output: - [ - { - "head": "SAM ALTMAN", - "head_type": "PERSON", - "head_description": "CEO of OpenAI, technology entrepreneur who announced the release of GPT-4 at the San Francisco press conference", - "relation": "Sam Altman is the CEO of OpenAI and announced the release of GPT-4", - "relation_strength": 9, - "tail": "OPENAI", - "tail_type": "ORGANIZATION", - "tail_description": "Artificial intelligence research company that developed GPT-4, a major AI milestone released on March 14, 2023" - }, - { - "head": "MICROSOFT", - "head_type": "ORGANIZATION", - "head_description": "Technology corporation and strategic partner of OpenAI, integrated GPT-4 into Bing search engine and Office suite", - "relation": "Microsoft is a strategic partner of OpenAI and integrated GPT-4 into its products", - "relation_strength": 8, - "tail": "OPENAI", - "tail_type": "ORGANIZATION", - "tail_description": "Artificial intelligence research company that developed GPT-4, partnered with Microsoft" - } - ] - - Example 2: - Text: - 2023年9月,阿里巴巴集团在杭州云栖大会上发布了通义千问2.0大模型。阿里云智能集团CEO张勇表示,通义千问将全面接入阿里巴巴旗下所有产品。同时,阿里巴巴宣布开源通义千问70亿参数模型,供开发者免费使用。 - Output: - [ - { - "head": "张勇", - "head_type": "PERSON", - "head_description": "阿里云智能集团CEO,负责发布通义千问2.0大模型,宣布通义千问全面接入阿里巴巴产品", - "relation": "张勇是阿里云智能集团的CEO,负责发布通义千问2.0", - "relation_strength": 9, - "tail": "阿里云智能集团", - "tail_type": "ORGANIZATION", - "tail_description": "阿里巴巴旗下云计算子公司,负责通义千问大模型的研发和发布" - }, - { - "head": "阿里巴巴集团", - "head_type": "ORGANIZATION", - "head_description": "中国互联网科技巨头,在2023年9月杭州云栖大会发布通义千问2.0,并开源70亿参数模型", - "relation": "阿里巴巴集团在杭州云栖大会上发布了通义千问2.0大模型", - "relation_strength": 8, - "tail": "杭州", - "tail_type": "LOCATION", - "tail_description": "中国浙江省省会城市,2023年云栖大会举办地" - } - ] - - Example 3: - Text: - 李明,著名企业家,于2025年1月1日在北京成立了"创新科技公司"。公司的主营业务是人工智能解决方案,并在成立当月获得了王芳女士的千万级天使投资。 - Output: - [ - { - "head": "李明", - "head_type": "PERSON", - "head_description": "著名企业家,创新科技公司创始人,于2025年1月1日在北京创办公司", - "relation": "李明是创新科技公司的创始人,于2025年1月1日成立了该公司", - "relation_strength": 8, - "tail": "创新科技公司", - "tail_type": "ORGANIZATION", - "tail_description": "人工智能解决方案公司,2025年1月在北京成立,获得千万级天使投资" - }, - { - "head": "王芳", - "head_type": "PERSON", - "head_description": "天使投资人,为创新科技公司提供千万级天使投资", - "relation": "王芳女士为创新科技公司提供了千万级的天使投资", - "relation_strength": 9, - "tail": "创新科技公司", - "tail_type": "ORGANIZATION", - "tail_description": "人工智能解决方案公司,成立当月即获得王芳女士千万级天使投资" - }, - { - "head": "创新科技公司", - "head_type": "ORGANIZATION", - "head_description": "人工智能解决方案公司,由著名企业家李明于2025年1月在北京创办", - "relation": "创新科技公司的注册地和成立地点是北京", - "relation_strength": 6, - "tail": "北京", - "tail_type": "LOCATION", - "tail_description": "中国首都,创新科技公司注册和成立地点" - } - ] - """; - - /** - * 创建配置好的 GraphTransformer - * @param chatModel LangChain4j ChatModel 实例 - * @return 配置好的 transformer - */ - public static LLMGraphTransformer create(ChatModel chatModel) { - - List allowedNodes = Arrays.asList(GRAPH_ENTITY_EXTRACTION_ENTITY_TYPES); - - return LLMGraphTransformer.builder() - .model(chatModel) - // 这会自动生成 "The 'head_type' and 'tail、_type' must be one of..." 的系统提示 - .allowedNodes(allowedNodes) - // 将 GraphRAG 的逻辑步骤注入到 System/User prompt 中 - .additionalInstructions(GRAPH_RAG_INSTRUCTIONS) - // 传入转换后的 JSON 格式示例 - .examples(GRAPH_RAG_JSON_EXAMPLES) - // 建议设置重试,因为 JSON 格式偶尔可能出错 - .maxAttempts(2) - .build(); - } -} \ No newline at end of file +package com.wemirr.platform.ai.core.provider.graph; + +import com.fasterxml.jackson.databind.JsonNode; +import com.fasterxml.jackson.databind.ObjectMapper; +import org.springframework.ai.chat.messages.UserMessage; +import org.springframework.ai.chat.model.ChatModel; +import org.springframework.ai.chat.prompt.Prompt; +import org.springframework.ai.document.Document; + +import java.util.Arrays; +import java.util.LinkedHashSet; +import java.util.List; +import java.util.Map; +import java.util.Set; +import java.util.regex.Matcher; +import java.util.regex.Pattern; + +/** + * Spring AI 图谱抽取器工厂。 + * + * @author xJh + * @since 2025/12/12 + */ +public class GraphRagTransformerFactory { + + public static final String[] GRAPH_ENTITY_EXTRACTION_ENTITY_TYPES = { + "ORGANIZATION", "PERSON", "LOCATION", "EVENT", "DATE" + }; + + private static final ObjectMapper OBJECT_MAPPER = new ObjectMapper(); + private static final Pattern JSON_ARRAY_PATTERN = Pattern.compile("\\[.*]", Pattern.DOTALL); + + private static final String GRAPH_RAG_INSTRUCTIONS = """ + -Goal- + Given a text document and a list of entity types, identify all entities of those types and all relationships among the identified entities. + + -Steps- + 1. Identify entities. For each entity extract entity_name, entity_type, entity_description. + 2. Identify clearly related source_entity and target_entity pairs. + 3. Return a strict JSON list. Do not output markdown or explanation. + 4. JSON object schema: + { + "head": "source_entity name", + "head_type": "source_entity type", + "head_description": "source_entity description", + "relation": "relationship description", + "relation_strength": 8, + "tail": "target_entity name", + "tail_type": "target_entity type", + "tail_description": "target_entity description" + } + 5. Allowed entity types: %s + """; + + public static SpringAiGraphTransformer create(ChatModel chatModel) { + return new SpringAiGraphTransformer(chatModel); + } + + public record SpringAiGraphTransformer(ChatModel chatModel) { + + public List transformAll(List documents) { + return documents.stream().map(this::transform).toList(); + } + + private GraphDocument transform(Document document) { + String prompt = GRAPH_RAG_INSTRUCTIONS.formatted(Arrays.toString(GRAPH_ENTITY_EXTRACTION_ENTITY_TYPES)) + + "\n\nText:\n" + document.getText(); + String response = chatModel.call(new Prompt(new UserMessage(prompt))) + .getResult().getOutput().getText(); + return parseGraphDocument(document, response); + } + + private GraphDocument parseGraphDocument(Document source, String response) { + try { + String json = extractJsonArray(response); + JsonNode array = OBJECT_MAPPER.readTree(json); + if (!array.isArray()) { + throw new IllegalArgumentException("GraphRAG 响应不是 JSON 数组"); + } + Set nodes = new LinkedHashSet<>(); + Set edges = new LinkedHashSet<>(); + for (JsonNode item : array) { + GraphNode head = node(item, "head", "head_type", "head_description"); + GraphNode tail = node(item, "tail", "tail_type", "tail_description"); + String relation = text(item, "relation", "RELATED_TO"); + nodes.add(head); + nodes.add(tail); + edges.add(new GraphEdge(head, tail, relation, Map.of( + "description", relation, + "strength", text(item, "relation_strength", "") + ))); + } + return new GraphDocument(nodes, edges, source); + } catch (Exception e) { + throw new IllegalStateException("GraphRAG JSON 抽取解析失败: " + response, e); + } + } + + private GraphNode node(JsonNode item, String nameField, String typeField, String descriptionField) { + String id = text(item, nameField, "UNKNOWN"); + String type = text(item, typeField, "UNKNOWN"); + return new GraphNode(id, type, Map.of("description", text(item, descriptionField, ""))); + } + + private String extractJsonArray(String response) { + if (response == null || response.isBlank()) { + throw new IllegalArgumentException("GraphRAG 响应为空"); + } + Matcher matcher = JSON_ARRAY_PATTERN.matcher(response); + if (!matcher.find()) { + throw new IllegalArgumentException("GraphRAG 响应不包含 JSON 数组"); + } + return matcher.group(); + } + + private String text(JsonNode node, String field, String defaultValue) { + JsonNode value = node.get(field); + if (value == null || value.isNull()) { + return defaultValue; + } + return value.asText(defaultValue); + } + } +} diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/graph/GraphRetriever.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/graph/GraphRetriever.java index 7be4b611..149b8fc0 100644 --- a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/graph/GraphRetriever.java +++ b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/graph/GraphRetriever.java @@ -1,7 +1,5 @@ package com.wemirr.platform.ai.core.provider.graph; -import dev.langchain4j.rag.content.Content; - import java.util.List; import static com.wemirr.platform.ai.core.constant.AiServiceConstants.*; @@ -99,22 +97,6 @@ public interface GraphRetriever { double scoreThreshold, int searchLimit, int hopDepth, int maxTriples); - // ==================== Content 转换 ==================== - - /** - * 将三元组列表转换为 Content 对象 - * - * @param triples 三元组列表 - * @return Content 列表 - */ - default List toContents(List triples) { - if (triples == null || triples.isEmpty()) { - return List.of(); - } - String context = "知识图谱检索结果:\n" + String.join("\n", triples); - return List.of(Content.from(context)); - } - // ==================== Prompt 构建 ==================== /** diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/graph/GraphStore.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/graph/GraphStore.java index 71b80b13..fb25d2ae 100644 --- a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/graph/GraphStore.java +++ b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/graph/GraphStore.java @@ -1,7 +1,5 @@ package com.wemirr.platform.ai.core.provider.graph; -import dev.langchain4j.community.data.document.graph.GraphDocument; - import java.util.List; import java.util.Map; diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/graph/neo4j/Neo4jGraphRetriever.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/graph/neo4j/Neo4jGraphRetriever.java index 11ee9e0d..14511a7a 100644 --- a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/graph/neo4j/Neo4jGraphRetriever.java +++ b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/graph/neo4j/Neo4jGraphRetriever.java @@ -6,14 +6,12 @@ import com.wemirr.platform.ai.domain.entity.KnowledgeBase; import com.wemirr.platform.ai.domain.entity.ModelEntity; import com.wemirr.platform.ai.service.KnowledgeBaseService; import com.wemirr.platform.ai.service.ModelService; -import dev.langchain4j.data.embedding.Embedding; -import dev.langchain4j.model.embedding.EmbeddingModel; -import dev.langchain4j.model.output.Response; import lombok.RequiredArgsConstructor; import lombok.extern.slf4j.Slf4j; import org.neo4j.driver.Record; import org.neo4j.driver.Session; import org.neo4j.driver.Values; +import org.springframework.ai.embedding.EmbeddingModel; import org.springframework.boot.autoconfigure.condition.ConditionalOnProperty; import org.springframework.stereotype.Component; @@ -63,8 +61,11 @@ public class Neo4jGraphRetriever implements GraphRetriever { try { // 1. 将用户问题转化为向量 - Response embeddingResponse = embeddingModel.embed(question); - List queryVector = embeddingResponse.content().vectorAsList(); + float[] embedding = embeddingModel.embed(question); + List queryVector = new ArrayList<>(embedding.length); + for (float value : embedding) { + queryVector.add(value); + } // 2. 使用向量索引查询最近邻节点 try (Session session = graphStore.getDriver().session()) { diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/graph/neo4j/Neo4jGraphStore.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/graph/neo4j/Neo4jGraphStore.java index 9e874f10..1d653e56 100644 --- a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/graph/neo4j/Neo4jGraphStore.java +++ b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/graph/neo4j/Neo4jGraphStore.java @@ -1,13 +1,13 @@ package com.wemirr.platform.ai.core.provider.graph.neo4j; import com.wemirr.platform.ai.core.config.VectorStoreProperties; +import com.wemirr.platform.ai.core.provider.graph.GraphDocument; import com.wemirr.platform.ai.core.provider.graph.GraphStore; -import dev.langchain4j.community.data.document.graph.GraphDocument; -import dev.langchain4j.model.embedding.EmbeddingModel; import jakarta.annotation.PreDestroy; import lombok.Getter; import lombok.extern.slf4j.Slf4j; import org.neo4j.driver.*; +import org.springframework.ai.embedding.EmbeddingModel; import org.springframework.boot.autoconfigure.condition.ConditionalOnProperty; import org.springframework.stereotype.Component; diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/graph/neo4j/Neo4jGraphWriter.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/graph/neo4j/Neo4jGraphWriter.java index bd2a56c5..2b381e5c 100644 --- a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/graph/neo4j/Neo4jGraphWriter.java +++ b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/graph/neo4j/Neo4jGraphWriter.java @@ -1,16 +1,14 @@ package com.wemirr.platform.ai.core.provider.graph.neo4j; -import dev.langchain4j.community.data.document.graph.GraphDocument; -import dev.langchain4j.community.data.document.graph.GraphEdge; -import dev.langchain4j.community.data.document.graph.GraphNode; -import dev.langchain4j.data.embedding.Embedding; -import dev.langchain4j.model.embedding.EmbeddingModel; -import dev.langchain4j.model.output.Response; +import com.wemirr.platform.ai.core.provider.graph.GraphDocument; +import com.wemirr.platform.ai.core.provider.graph.GraphEdge; +import com.wemirr.platform.ai.core.provider.graph.GraphNode; import lombok.Builder; import lombok.extern.slf4j.Slf4j; import org.neo4j.driver.Driver; import org.neo4j.driver.Session; import org.neo4j.driver.Values; +import org.springframework.ai.embedding.EmbeddingModel; import java.util.List; import java.util.Map; @@ -174,8 +172,12 @@ public class Neo4jGraphWriter { return null; } try { - Response response = embeddingModel.embed(text); - return response.content().vectorAsList(); + float[] vector = embeddingModel.embed(text); + List result = new java.util.ArrayList<>(vector.length); + for (float value : vector) { + result.add(value); + } + return result; } catch (Exception e) { log.warn("生成向量嵌入失败: text={}, error={}", text.substring(0, Math.min(50, text.length())), e.getMessage()); return null; @@ -216,7 +218,7 @@ public class Neo4jGraphWriter { * 创建源文档节点并关联 */ private void createSourceDocument(Session session, GraphDocument doc) { - String docText = doc.source().text(); + String docText = doc.source().getText(); String docId = "doc_" + docText.hashCode(); String escapedLabel = label != null ? escapeLabel(label) : null; diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/mcp/ContextAwareMcpToolProvider.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/mcp/ContextAwareMcpToolProvider.java deleted file mode 100644 index 656a7f9a..00000000 --- a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/mcp/ContextAwareMcpToolProvider.java +++ /dev/null @@ -1,89 +0,0 @@ -package com.wemirr.platform.ai.core.provider.mcp; - -import com.wemirr.platform.ai.service.McpConnectionManager; -import dev.langchain4j.agent.tool.ToolSpecification; -import dev.langchain4j.mcp.McpToolExecutor; -import dev.langchain4j.mcp.client.McpClient; -import dev.langchain4j.service.tool.ToolExecutor; -import dev.langchain4j.service.tool.ToolProvider; -import dev.langchain4j.service.tool.ToolProviderRequest; -import dev.langchain4j.service.tool.ToolProviderResult; -import lombok.extern.slf4j.Slf4j; - -import java.util.List; - -/** - * 上下文感知的 MCP 工具提供者 - *

    - * 通过构造函数注入上下文,避免 ThreadLocal 在异步场景下的问题。 - * 每个智能体会话创建独立的实例。 - * - * @author xJh - * @since 2025/12/28 - */ -@Slf4j -public class ContextAwareMcpToolProvider implements ToolProvider { - - private final McpConnectionManager mcpConnectionManager; - private final McpToolProviderContext context; - - /** - * 构造函数 - * - * @param mcpConnectionManager MCP 连接管理器 - * @param context MCP 工具上下文 - */ - public ContextAwareMcpToolProvider(McpConnectionManager mcpConnectionManager, McpToolProviderContext context) { - this.mcpConnectionManager = mcpConnectionManager; - this.context = context; - } - - @Override - public ToolProviderResult provideTools(ToolProviderRequest request) { - // 检查上下文是否有效 - if (context == null || !context.isValid()) { - log.debug("No MCP servers configured for current context"); - return ToolProviderResult.builder().build(); - } - - List mcpServerIds = context.getMcpServerIds(); - ToolProviderResult.Builder builder = ToolProviderResult.builder(); - int totalTools = 0; - - // 加载配置的 MCP 服务器工具 - for (Long configId : mcpServerIds) { - try { - McpClient client = mcpConnectionManager.getClient(configId); - List tools = client.listTools(); - - if (tools != null) { - for (ToolSpecification originalSpec : tools) { - String originalToolName = originalSpec.name(); - - // 为避免工具名称冲突,添加前缀 - String uniqueToolName = "mcp_" + configId + "_" + originalToolName; - - // 创建新的 ToolSpecification,使用唯一名称 - ToolSpecification newSpec = originalSpec.toBuilder() - .name(uniqueToolName) - .build(); - - // 使用官方的 McpToolExecutor,传入原始工具名称 - ToolExecutor toolExecutor = new McpToolExecutor(client, originalToolName); - - // 添加工具和执行器 - builder.add(newSpec, toolExecutor); - totalTools++; - } - } - } catch (Exception e) { - log.warn("Failed to provide tools for MCP server ID: {}, session: {}", - configId, context.getSessionId(), e); - } - } - - log.info("Loaded {} MCP tools from {} servers for agent: {}, session: {}", - totalTools, mcpServerIds.size(), context.getAgentId(), context.getSessionId()); - return builder.build(); - } -} diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/mcp/DynamicMcpToolProvider.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/mcp/DynamicMcpToolProvider.java deleted file mode 100644 index 9331762b..00000000 --- a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/mcp/DynamicMcpToolProvider.java +++ /dev/null @@ -1,95 +0,0 @@ -package com.wemirr.platform.ai.core.provider.mcp; - -import com.wemirr.platform.ai.service.McpConnectionManager; -import dev.langchain4j.agent.tool.ToolSpecification; -import dev.langchain4j.mcp.McpToolExecutor; -import dev.langchain4j.mcp.client.McpClient; -import dev.langchain4j.service.tool.ToolExecutor; -import dev.langchain4j.service.tool.ToolProvider; -import dev.langchain4j.service.tool.ToolProviderRequest; -import dev.langchain4j.service.tool.ToolProviderResult; -import lombok.RequiredArgsConstructor; -import lombok.extern.slf4j.Slf4j; -import org.springframework.stereotype.Service; - -import java.util.List; - -/** - * 动态MCP工具提供者 - * - * @author xJh - * @since 2025/12/07 - */ -@Slf4j -@Service -@RequiredArgsConstructor -public class DynamicMcpToolProvider implements ToolProvider { - - private final McpConnectionManager mcpConnectionManager; - - // 存储当前智能体的 MCP 配置 ID 列表 - private ThreadLocal> agentMcpServerIds = new ThreadLocal<>(); - - /** - * 设置当前智能体的 MCP 服务器配置 - */ - public void setAgentMcpServerIds(List mcpServerIds) { - this.agentMcpServerIds.set(mcpServerIds); - } - - /** - * 清除当前线程的 MCP 配置 - */ - public void clearAgentMcpServerIds() { - this.agentMcpServerIds.remove(); - } - - @Override - public ToolProviderResult provideTools(ToolProviderRequest request) { - // 获取当前智能体配置的 MCP Server ID 列表 - List mcpServerIds = agentMcpServerIds.get(); - - // 如果智能体没有配置 MCP,返回空工具 - if (mcpServerIds == null || mcpServerIds.isEmpty()) { - log.debug("No MCP servers configured for current agent"); - return ToolProviderResult.builder().build(); - } - - ToolProviderResult.Builder builder = ToolProviderResult.builder(); - int totalTools = 0; - - // 只加载智能体配置的 MCP 服务器 - for (Long configId : mcpServerIds) { - try { - McpClient client = mcpConnectionManager.getClient(configId); - List tools = client.listTools(); - - if (tools != null) { - for (ToolSpecification originalSpec : tools) { - String originalToolName = originalSpec.name(); - - // 为避免工具名称冲突,添加前缀 - String uniqueToolName = "mcp_" + configId + "_" + originalToolName; - - // 创建新的 ToolSpecification,使用唯一名称 - ToolSpecification newSpec = originalSpec.toBuilder() - .name(uniqueToolName) - .build(); - - // 使用官方的 McpToolExecutor,传入原始工具名称 - ToolExecutor toolExecutor = new McpToolExecutor(client, originalToolName); - - // 添加工具和执行器 - builder.add(newSpec, toolExecutor); - totalTools++; - } - } - } catch (Exception e) { - log.warn("Failed to provide tools for MCP server ID: {}", configId, e); - } - } - - log.info("Loaded {} MCP tools from {} servers", totalTools, mcpServerIds.size()); - return builder.build(); - } -} diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/mcp/McpClientHandle.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/mcp/McpClientHandle.java new file mode 100644 index 00000000..7eee397b --- /dev/null +++ b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/mcp/McpClientHandle.java @@ -0,0 +1,20 @@ +package com.wemirr.platform.ai.core.provider.mcp; + +import java.util.List; + +/** + * MCP 客户端句柄。 + * + * @author xJh + * @since 2026/06/27 + */ +public interface McpClientHandle extends AutoCloseable { + + List listTools(); + + @Override + void close(); + + record McpToolDescriptor(String name, String description) { + } +} diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/mcp/McpToolProviderFactory.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/mcp/McpToolProviderFactory.java deleted file mode 100644 index 284cd180..00000000 --- a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/mcp/McpToolProviderFactory.java +++ /dev/null @@ -1,60 +0,0 @@ -package com.wemirr.platform.ai.core.provider.mcp; - -import com.wemirr.platform.ai.service.McpConnectionManager; -import dev.langchain4j.service.tool.ToolProvider; -import lombok.RequiredArgsConstructor; -import lombok.extern.slf4j.Slf4j; -import org.springframework.stereotype.Component; - -import java.util.List; - -/** - * MCP 工具提供者工厂 - *

    - * 负责创建上下文感知的 MCP 工具提供者实例 - * - * @author xJh - * @since 2025/12/28 - */ -@Slf4j -@Component -@RequiredArgsConstructor -public class McpToolProviderFactory { - - private final McpConnectionManager mcpConnectionManager; - - /** - * 为智能体创建 MCP 工具提供者 - * - * @param agentId 智能体 ID - * @param mcpServerIds MCP 服务器配置 ID 列表 - * @param sessionId 会话 ID - * @return MCP 工具提供者 - */ - public ToolProvider create(Long agentId, List mcpServerIds, String sessionId) { - if (mcpServerIds == null || mcpServerIds.isEmpty()) { - log.debug("No MCP servers configured for agent: {}", agentId); - return null; - } - - McpToolProviderContext context = McpToolProviderContext.builder() - .agentId(agentId) - .mcpServerIds(mcpServerIds) - .sessionId(sessionId) - .build(); - - log.debug("Creating MCP tool provider for agent: {}, servers: {}", agentId, mcpServerIds); - return new ContextAwareMcpToolProvider(mcpConnectionManager, context); - } - - /** - * 为智能体创建 MCP 工具提供者(无会话 ID) - * - * @param agentId 智能体 ID - * @param mcpServerIds MCP 服务器配置 ID 列表 - * @return MCP 工具提供者 - */ - public ToolProvider create(Long agentId, List mcpServerIds) { - return create(agentId, mcpServerIds, null); - } -} diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/retrieval/ContentRetrieverRegistry.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/retrieval/ContentRetrieverRegistry.java deleted file mode 100644 index 9c69ca0a..00000000 --- a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/retrieval/ContentRetrieverRegistry.java +++ /dev/null @@ -1,86 +0,0 @@ -package com.wemirr.platform.ai.core.provider.retrieval; - -import com.wemirr.platform.ai.core.assistant.service.RagAssistantParams; -import dev.langchain4j.model.chat.ChatModel; -import dev.langchain4j.rag.content.retriever.ContentRetriever; -import lombok.extern.slf4j.Slf4j; -import org.springframework.stereotype.Component; - -import java.util.LinkedHashMap; -import java.util.List; -import java.util.Map; -import java.util.Optional; - -/** - * 内容检索器注册中心 - *

    - * 管理所有检索策略,根据参数配置创建对应的检索器 - * - * @author xJh - * @since 2025/12/28 - */ -@Slf4j -@Component -public class ContentRetrieverRegistry { - - private final List strategies; - - public ContentRetrieverRegistry(List strategies) { - this.strategies = strategies; - log.info("注册了 {} 个内容检索策略: {}", strategies.size(), - strategies.stream().map(ContentRetrieverStrategy::getType).toList()); - } - - /** - * 根据参数创建所有支持的检索器 - * - * @param params RAG 参数 - * @param chatModel 聊天模型 - * @return 检索器及其描述的映射 - */ - public Map createRetrievers(RagAssistantParams params, ChatModel chatModel) { - Map retrieverMap = new LinkedHashMap<>(); - - for (ContentRetrieverStrategy strategy : strategies) { - if (strategy.supports(params)) { - strategy.createRetriever(params, chatModel) - .ifPresent(retriever -> { - retrieverMap.put(retriever, strategy.getDescription()); - log.debug("启用检索策略: type={}, description={}", - strategy.getType(), strategy.getDescription()); - }); - } - } - - if (retrieverMap.isEmpty()) { - log.warn("没有可用的检索器,请检查参数配置"); - } else { - log.info("创建了 {} 个检索器", retrieverMap.size()); - } - - return retrieverMap; - } - - /** - * 获取所有已注册的策略类型 - * - * @return 策略类型列表 - */ - public List getRegisteredTypes() { - return strategies.stream() - .map(ContentRetrieverStrategy::getType) - .toList(); - } - - /** - * 根据类型获取策略 - * - * @param type 策略类型 - * @return 策略实例 - */ - public Optional getStrategy(String type) { - return strategies.stream() - .filter(s -> s.getType().equals(type)) - .findFirst(); - } -} diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/retrieval/ContentRetrieverStrategy.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/retrieval/ContentRetrieverStrategy.java deleted file mode 100644 index 68a866b2..00000000 --- a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/retrieval/ContentRetrieverStrategy.java +++ /dev/null @@ -1,49 +0,0 @@ -package com.wemirr.platform.ai.core.provider.retrieval; - -import com.wemirr.platform.ai.core.assistant.service.RagAssistantParams; -import dev.langchain4j.model.chat.ChatModel; -import dev.langchain4j.rag.content.retriever.ContentRetriever; - -import java.util.Optional; - -/** - * 内容检索策略接口 - *

    - * 定义不同类型检索器的创建策略,支持向量检索、图谱检索等 - * - * @author xJh - * @since 2025/12/28 - */ -public interface ContentRetrieverStrategy { - - /** - * 获取策略类型 - * - * @return 策略类型标识 - */ - String getType(); - - /** - * 获取检索器描述 - * - * @return 检索器描述 - */ - String getDescription(); - - /** - * 判断是否支持当前参数配置 - * - * @param params RAG 参数 - * @return 是否支持 - */ - boolean supports(RagAssistantParams params); - - /** - * 创建内容检索器 - * - * @param params RAG 参数 - * @param chatModel 聊天模型(部分检索器可能需要) - * @return 内容检索器,如果创建失败返回空 - */ - Optional createRetriever(RagAssistantParams params, ChatModel chatModel); -} diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/retrieval/GraphContentRetrieverStrategy.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/retrieval/GraphContentRetrieverStrategy.java deleted file mode 100644 index 2b43babe..00000000 --- a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/retrieval/GraphContentRetrieverStrategy.java +++ /dev/null @@ -1,80 +0,0 @@ -package com.wemirr.platform.ai.core.provider.retrieval; - -import com.wemirr.platform.ai.core.assistant.service.RagAssistantParams; -import com.wemirr.platform.ai.core.provider.graph.GraphContentRetriever; -import com.wemirr.platform.ai.core.provider.graph.GraphRagService; -import dev.langchain4j.model.chat.ChatModel; -import dev.langchain4j.rag.content.retriever.ContentRetriever; -import lombok.extern.slf4j.Slf4j; -import org.springframework.beans.factory.annotation.Autowired; -import org.springframework.stereotype.Component; - -import java.util.Optional; - -/** - * 图谱检索策略 - *

    - * 基于知识图谱的内容检索实现 - * - * @author xJh - * @since 2025/12/28 - */ -@Slf4j -@Component -public class GraphContentRetrieverStrategy implements ContentRetrieverStrategy { - - /** - * GraphRAG 服务(可选,仅在启用图谱功能时注入) - */ - @Autowired(required = false) - private GraphRagService graphRagService; - - @Override - public String getType() { - return "GRAPH"; - } - - @Override - public String getDescription() { - return "知识图谱(实体关系、结构化数据)"; - } - - @Override - public boolean supports(RagAssistantParams params) { - return Boolean.TRUE.equals(params.getEnableGraphRetrieval()) - && graphRagService != null - && params.getEffectiveGraphKbId() != null; - } - - @Override - public Optional createRetriever(RagAssistantParams params, ChatModel chatModel) { - if (graphRagService == null) { - log.warn("GraphRagService 未注入,无法创建图谱检索器"); - return Optional.empty(); - } - - try { - String graphKbId = params.getEffectiveGraphKbId(); - if (graphKbId == null) { - log.warn("图谱知识库ID为空"); - return Optional.empty(); - } - - ContentRetriever retriever = GraphContentRetriever.builder() - .graphRagService(graphRagService) - .chatModel(chatModel) - .knowledgeBaseId(graphKbId) - .maxResults(params.getGraphMaxResults()) - .silentOnEmpty(true) - .build(); - - log.debug("图谱检索器创建成功: graphKbId={}, maxResults={}", - graphKbId, params.getGraphMaxResults()); - - return Optional.of(retriever); - } catch (Exception e) { - log.error("创建图谱检索器失败: graphKbId={}", params.getEffectiveGraphKbId(), e); - return Optional.empty(); - } - } -} diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/retrieval/VectorContentRetrieverStrategy.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/retrieval/VectorContentRetrieverStrategy.java deleted file mode 100644 index 915b8626..00000000 --- a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/retrieval/VectorContentRetrieverStrategy.java +++ /dev/null @@ -1,84 +0,0 @@ -package com.wemirr.platform.ai.core.provider.retrieval; - -import com.wemirr.framework.ai.core.provider.embedding.EmbeddingModelRegistry; -import com.wemirr.platform.ai.core.assistant.service.RagAssistantParams; -import com.wemirr.platform.ai.core.provider.vector.VectorStoreFactory; -import com.wemirr.platform.ai.domain.entity.KnowledgeBase; -import com.wemirr.platform.ai.service.KnowledgeBaseService; -import dev.langchain4j.data.segment.TextSegment; -import dev.langchain4j.model.chat.ChatModel; -import dev.langchain4j.model.embedding.EmbeddingModel; -import dev.langchain4j.rag.content.retriever.ContentRetriever; -import dev.langchain4j.rag.content.retriever.EmbeddingStoreContentRetriever; -import dev.langchain4j.store.embedding.EmbeddingStore; -import lombok.RequiredArgsConstructor; -import lombok.extern.slf4j.Slf4j; -import org.springframework.stereotype.Component; - -import java.util.Optional; - -/** - * 向量检索策略 - *

    - * 基于向量相似度的内容检索实现 - * - * @author xJh - * @since 2025/12/28 - */ -@Slf4j -@Component -@RequiredArgsConstructor -public class VectorContentRetrieverStrategy implements ContentRetrieverStrategy { - - private final VectorStoreFactory vectorStoreFactory; - private final KnowledgeBaseService knowledgeBaseService; - private final EmbeddingModelRegistry embeddingModelRegistry; - - @Override - public String getType() { - return "VECTOR"; - } - - @Override - public String getDescription() { - return "内部知识库(文档、手册、策略等非结构化内容)"; - } - - @Override - public boolean supports(RagAssistantParams params) { - return Boolean.TRUE.equals(params.getEnableVectorRetrieval()) - && params.getEmbeddingModelEntity() != null - && params.getKbId() != null; - } - - @Override - public Optional createRetriever(RagAssistantParams params, ChatModel chatModel) { - try { - KnowledgeBase knowledgeBase = knowledgeBaseService.getById(params.getKbId()); - if (knowledgeBase == null) { - log.warn("知识库不存在: kbId={}", params.getKbId()); - return Optional.empty(); - } - - EmbeddingStore embeddingStore = vectorStoreFactory.createForKnowledgeBase( - knowledgeBase, params.getEmbeddingModelEntity()); - EmbeddingModel embeddingModel = embeddingModelRegistry.getFactory(params.getEmbeddingModelEntity()) - .createModel(params.getEmbeddingModelEntity()); - - ContentRetriever retriever = EmbeddingStoreContentRetriever.builder() - .embeddingStore(embeddingStore) - .embeddingModel(embeddingModel) - .maxResults(params.getMaxResults()) - .minScore(params.getMinScore()) - .build(); - - log.debug("向量检索器创建成功: kbId={}, maxResults={}, minScore={}", - params.getKbId(), params.getMaxResults(), params.getMinScore()); - - return Optional.of(retriever); - } catch (Exception e) { - log.error("创建向量检索器失败: kbId={}", params.getKbId(), e); - return Optional.empty(); - } - } -} diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/scoring/CohereScoringModelProvider.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/scoring/CohereScoringModelProvider.java deleted file mode 100644 index abb27d60..00000000 --- a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/scoring/CohereScoringModelProvider.java +++ /dev/null @@ -1,61 +0,0 @@ -package com.wemirr.platform.ai.core.provider.scoring; - -import com.wemirr.platform.ai.domain.entity.ModelEntity; -import dev.langchain4j.model.cohere.CohereScoringModel; -import dev.langchain4j.model.scoring.ScoringModel; -import org.springframework.stereotype.Component; - -import java.time.Duration; -import java.util.Map; - -/** - * Cohere 重排序模型提供者 - *

    - * Cohere 提供高质量的多语言重排序模型 - * - * @author xJh - * @since 2025/12/18 - */ -@Component -public class CohereScoringModelProvider implements ScoringModelProvider { - - private static final String DEFAULT_MODEL = "rerank-multilingual-v3.0"; - - @Override - public boolean supports(ModelEntity config) { - String provider = config.getProvider() != null ? config.getProvider().getValue() : null; - return "cohere".equalsIgnoreCase(provider); - } - - @Override - public ScoringModel createModel(ModelEntity config) { - String modelName = config.getName() != null ? config.getName() : DEFAULT_MODEL; - - CohereScoringModel.CohereScoringModelBuilder builder = CohereScoringModel.builder() - .apiKey(config.getApiKey()) - .modelName(modelName); - - // 从 variables 读取额外配置 - Map variables = config.getVariables(); - if (variables != null) { - if (variables.containsKey("timeout")) { - builder.timeout(Duration.ofSeconds(((Number) variables.get("timeout")).longValue())); - } - if (variables.containsKey("maxRetries")) { - builder.maxRetries(((Number) variables.get("maxRetries")).intValue()); - } - } - - // 支持自定义 baseUrl(用于代理或私有化部署) - if (config.getBaseUrl() != null && !config.getBaseUrl().isBlank()) { - builder.baseUrl(config.getBaseUrl()); - } - - return builder.build(); - } - - @Override - public String providerId() { - return "cohere"; - } -} diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/scoring/JinaScoringModelProvider.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/scoring/JinaScoringModelProvider.java deleted file mode 100644 index 41e034fa..00000000 --- a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/scoring/JinaScoringModelProvider.java +++ /dev/null @@ -1,54 +0,0 @@ -package com.wemirr.platform.ai.core.provider.scoring; - -import com.wemirr.platform.ai.domain.entity.ModelEntity; -import dev.langchain4j.model.jina.JinaScoringModel; -import dev.langchain4j.model.scoring.ScoringModel; -import org.springframework.stereotype.Component; - -import java.time.Duration; -import java.util.Map; - -/** - * Jina 重排序模型提供者 - * - * @author xJh - * @since 2025/12/18 - */ -@Component -public class JinaScoringModelProvider implements ScoringModelProvider { - - private static final String DEFAULT_MODEL = "jina-reranker-v2-base-multilingual"; - - @Override - public boolean supports(ModelEntity config) { - String provider = config.getProvider() != null ? config.getProvider().getValue() : null; - return "jina".equalsIgnoreCase(provider) || "jina-ai".equalsIgnoreCase(provider); - } - - @Override - public ScoringModel createModel(ModelEntity config) { - String modelName = config.getName() != null ? config.getName() : DEFAULT_MODEL; - - JinaScoringModel.JinaScoringModelBuilder builder = JinaScoringModel.builder() - .apiKey(config.getApiKey()) - .modelName(modelName); - - // 从 variables 读取额外配置 - Map variables = config.getVariables(); - if (variables != null) { - if (variables.containsKey("timeout")) { - builder.timeout(Duration.ofSeconds(((Number) variables.get("timeout")).longValue())); - } - if (variables.containsKey("maxRetries")) { - builder.maxRetries(((Number) variables.get("maxRetries")).intValue()); - } - } - - return builder.build(); - } - - @Override - public String providerId() { - return "jina"; - } -} diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/scoring/ScoringModelProvider.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/scoring/ScoringModelProvider.java deleted file mode 100644 index 326643f5..00000000 --- a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/scoring/ScoringModelProvider.java +++ /dev/null @@ -1,38 +0,0 @@ -package com.wemirr.platform.ai.core.provider.scoring; - -import com.wemirr.platform.ai.domain.entity.ModelEntity; -import dev.langchain4j.model.scoring.ScoringModel; - -/** - * 重排序模型提供者接口 - *

    - * 定义重排序模型的通用抽象,支持多种重排序服务(Jina、Cohere、OpenAI 等) - * - * @author xJh - * @since 2025/12/18 - */ -public interface ScoringModelProvider { - - /** - * 是否支持该配置 - * - * @param config 模型配置 - * @return 是否支持 - */ - boolean supports(ModelEntity config); - - /** - * 创建重排序模型实例 - * - * @param config 模型配置 - * @return ScoringModel 实例 - */ - ScoringModel createModel(ModelEntity config); - - /** - * 获取提供商标识 - * - * @return 提供商 ID - */ - String providerId(); -} diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/scoring/ScoringModelProviderRegistry.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/scoring/ScoringModelProviderRegistry.java deleted file mode 100644 index cebacd16..00000000 --- a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/scoring/ScoringModelProviderRegistry.java +++ /dev/null @@ -1,100 +0,0 @@ -package com.wemirr.platform.ai.core.provider.scoring; - -import com.wemirr.framework.ai.core.enums.AiProvider; -import com.wemirr.framework.commons.exception.CheckedException; -import com.wemirr.platform.ai.domain.entity.ModelEntity; -import jakarta.annotation.PostConstruct; -import lombok.extern.slf4j.Slf4j; -import org.springframework.stereotype.Component; - -import java.util.HashMap; -import java.util.List; -import java.util.Map; -import java.util.stream.Collectors; - -/** - * 重排序模型提供者注册中心 - *

    - * 管理所有重排序模型提供者的注册和查找 - * - * @author xJh - * @since 2025/12/18 - */ -@Slf4j -@Component -public class ScoringModelProviderRegistry { - - private final Map providerMap = new HashMap<>(); - private final List providers; - - public ScoringModelProviderRegistry(List providers) { - this.providers = providers; - } - - @PostConstruct - public void init() { - for (ScoringModelProvider provider : providers) { - if (providerMap.containsKey(provider.providerId())) { - throw new IllegalArgumentException( - String.format("重复注册重排序模型提供商: %s", provider.providerId()) - ); - } - providerMap.put(provider.providerId(), provider); - log.info("注册重排序模型提供者: {}", provider.providerId()); - } - log.info("重排序模型提供者注册完成,共 {} 个", providers.size()); - } - - /** - * 根据配置查找匹配的提供者 - * - * @param config 模型配置 - * @return 匹配的提供者 - */ - public ScoringModelProvider getProvider(ModelEntity config) { - if (config.getProvider() == null) { - throw CheckedException.badRequest("模型提供商不能为空"); - } - if (config.getName() == null) { - throw CheckedException.badRequest("模型名称不能为空"); - } - - for (ScoringModelProvider provider : providers) { - if (provider.supports(config)) { - return provider; - } - } - - throw CheckedException.badRequest( - String.format("未找到支持的重排序模型提供商: provider=%s, model=%s", - config.getProvider(), config.getName()) - ); - } - - /** - * 获取所有可用的提供商 ID - * - * @return 提供商 ID 列表 - */ - public List getAvailableProviderIds() { - return providers.stream() - .map(ScoringModelProvider::providerId) - .filter(id -> { - try { - return AiProvider.of(id).isEnabled(); - } catch (Exception e) { - return true; - } - }) - .collect(Collectors.toList()); - } - - /** - * 检查是否有可用的重排序模型提供者 - * - * @return 是否可用 - */ - public boolean hasAvailableProviders() { - return !providers.isEmpty(); - } -} diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/scoring/ScoringModelService.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/scoring/ScoringModelService.java index e01dae15..917f9ae3 100644 --- a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/scoring/ScoringModelService.java +++ b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/scoring/ScoringModelService.java @@ -1,110 +1,27 @@ package com.wemirr.platform.ai.core.provider.scoring; -import com.wemirr.platform.ai.domain.entity.ModelEntity; -import dev.langchain4j.model.scoring.ScoringModel; -import lombok.RequiredArgsConstructor; import lombok.extern.slf4j.Slf4j; import org.springframework.stereotype.Service; -import java.util.List; -import java.util.Map; -import java.util.concurrent.ConcurrentHashMap; - /** - * 重排序模型服务 - *

    - * 提供重排序模型的获取和缓存管理 + * 重排序模型服务。 * * @author xJh * @since 2025/12/18 */ @Slf4j @Service -@RequiredArgsConstructor public class ScoringModelService { - private final ScoringModelProviderRegistry registry; - - /** - * 模型缓存,避免重复创建 - */ - private final Map modelCache = new ConcurrentHashMap<>(); - - /** - * 根据配置获取 ScoringModel 实例 - * - * @param config 模型配置 - * @return ScoringModel 实例 - */ - public ScoringModel getModel(ModelEntity config) { - String cacheKey = buildCacheKey(config); - return modelCache.computeIfAbsent(cacheKey, key -> { - ScoringModelProvider provider = registry.getProvider(config); - log.debug("创建重排序模型: provider={}, model={}", config.getProvider(), config.getName()); - return provider.createModel(config); - }); - } - - /** - * 获取提供者 - * - * @param config 模型配置 - * @return 提供者 - */ - public ScoringModelProvider getProvider(ModelEntity config) { - return registry.getProvider(config); - } - - /** - * 检查提供商是否可用 - * - * @param providerId 提供商 ID - * @return 是否可用 - */ public boolean isProviderAvailable(String providerId) { - return registry.getAvailableProviderIds().contains(providerId.toLowerCase()); + return false; } - /** - * 获取所有可用的提供商 ID - * - * @return 提供商 ID 列表 - */ - public List getAvailableProviderIds() { - return registry.getAvailableProviderIds(); - } - - /** - * 检查是否有可用的重排序模型 - * - * @return 是否可用 - */ public boolean hasAvailableProviders() { - return registry.hasAvailableProviders(); + return false; } - /** - * 清除模型缓存 - */ public void clearCache() { - modelCache.clear(); - log.info("重排序模型缓存已清除"); - } - - /** - * 移除指定模型的缓存 - * - * @param config 模型配置 - */ - public void evictCache(ModelEntity config) { - String cacheKey = buildCacheKey(config); - modelCache.remove(cacheKey); - log.debug("移除重排序模型缓存: {}", cacheKey); - } - - private String buildCacheKey(ModelEntity config) { - return String.format("%s:%s:%s", config.getProvider(), config.getName(), - config.getBaseUrl() != null ? config.getBaseUrl() : "default" - ); + log.debug("Spring AI 重排序模型缓存清理完成"); } } diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/text/TextModelCache.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/text/TextModelCache.java index 9a35d246..f49b5a12 100644 --- a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/text/TextModelCache.java +++ b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/text/TextModelCache.java @@ -4,8 +4,8 @@ import cn.hutool.extra.spring.SpringUtil; import com.wemirr.framework.ai.core.provider.text.TextModelProvider; import com.wemirr.framework.ai.core.provider.text.TextModelProviderRegistry; import com.wemirr.platform.ai.domain.entity.ModelEntity; -import dev.langchain4j.model.chat.ChatModel; -import dev.langchain4j.model.chat.StreamingChatModel; +import org.springframework.ai.chat.model.ChatModel; +import org.springframework.ai.chat.model.StreamingChatModel; import org.springframework.cache.annotation.Cacheable; import org.springframework.stereotype.Component; diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/text/TextModelService.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/text/TextModelService.java index 1afe6e36..659538e4 100644 --- a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/text/TextModelService.java +++ b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/text/TextModelService.java @@ -3,9 +3,9 @@ package com.wemirr.platform.ai.core.provider.text; import com.wemirr.framework.ai.core.provider.text.TextModelProvider; import com.wemirr.framework.ai.core.provider.text.TextModelProviderRegistry; import com.wemirr.platform.ai.domain.entity.ModelEntity; -import dev.langchain4j.model.chat.ChatModel; -import dev.langchain4j.model.chat.StreamingChatModel; import lombok.RequiredArgsConstructor; +import org.springframework.ai.chat.model.ChatModel; +import org.springframework.ai.chat.model.StreamingChatModel; import org.springframework.stereotype.Service; /** diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/vector/VectorStoreFactory.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/vector/VectorStoreFactory.java index f97ba387..4f82185d 100644 --- a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/vector/VectorStoreFactory.java +++ b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/provider/vector/VectorStoreFactory.java @@ -1,17 +1,22 @@ package com.wemirr.platform.ai.core.provider.vector; import com.wemirr.platform.ai.core.config.VectorStoreProperties; +import com.wemirr.platform.ai.core.provider.embedding.EmbeddingModelService; import com.wemirr.platform.ai.domain.entity.KnowledgeBase; import com.wemirr.platform.ai.domain.entity.ModelEntity; -import dev.langchain4j.data.segment.TextSegment; -import dev.langchain4j.store.embedding.EmbeddingStore; -import dev.langchain4j.store.embedding.inmemory.InMemoryEmbeddingStore; -import dev.langchain4j.store.embedding.milvus.MilvusEmbeddingStore; -import dev.langchain4j.store.embedding.pgvector.PgVectorEmbeddingStore; import io.milvus.client.MilvusServiceClient; import io.milvus.param.ConnectParam; +import io.milvus.param.IndexType; +import io.milvus.param.MetricType; import lombok.RequiredArgsConstructor; import lombok.extern.slf4j.Slf4j; +import org.springframework.ai.embedding.EmbeddingModel; +import org.springframework.ai.vectorstore.SimpleVectorStore; +import org.springframework.ai.vectorstore.VectorStore; +import org.springframework.ai.vectorstore.milvus.MilvusVectorStore; +import org.springframework.ai.vectorstore.pgvector.PgVectorStore; +import org.springframework.jdbc.core.JdbcTemplate; +import org.springframework.jdbc.datasource.DriverManagerDataSource; import org.springframework.stereotype.Component; import java.util.Map; @@ -36,11 +41,12 @@ import java.util.concurrent.ConcurrentHashMap; public class VectorStoreFactory { private final VectorStoreProperties properties; + private final EmbeddingModelService embeddingModelService; /** * 缓存已创建的向量存储实例 */ - private final Map> storeCache = new ConcurrentHashMap<>(); + private final Map storeCache = new ConcurrentHashMap<>(); /** * 根据知识库创建向量存储 @@ -49,7 +55,7 @@ public class VectorStoreFactory { * @param modelEntity 模型配置 * @return 向量存储实例 */ - public EmbeddingStore createForKnowledgeBase(KnowledgeBase knowledgeBase, ModelEntity modelEntity) { + public VectorStore createForKnowledgeBase(KnowledgeBase knowledgeBase, ModelEntity modelEntity) { String cacheKey = generateCacheKey(knowledgeBase, modelEntity); // 检查缓存 @@ -58,7 +64,7 @@ public class VectorStoreFactory { } // 创建新的向量存储实例 - EmbeddingStore store = createStore(knowledgeBase, modelEntity); + VectorStore store = createStore(knowledgeBase, modelEntity); storeCache.put(cacheKey, store); return store; @@ -69,31 +75,32 @@ public class VectorStoreFactory { * * @return 向量存储实例 */ - public EmbeddingStore createDefault() { + public VectorStore createDefault(ModelEntity modelEntity) { VectorStoreProperties.StoreType type = properties.getType(); + EmbeddingModel embeddingModel = embeddingModelService.getModel(modelEntity); return switch (type) { - case MILVUS -> createMilvus(properties.getMilvus().getCollectionName(), properties.getMilvus().getDimension()); - case PGVECTOR -> createPgVector(); - case IN_MEMORY -> new InMemoryEmbeddingStore<>(); + case MILVUS -> createMilvus(properties.getMilvus().getCollectionName(), properties.getMilvus().getDimension(), embeddingModel); + case PGVECTOR -> createPgVector(embeddingModel); + case IN_MEMORY -> SimpleVectorStore.builder(embeddingModel).build(); }; } /** * 创建向量存储实例 */ - private EmbeddingStore createStore(KnowledgeBase knowledgeBase, ModelEntity modelEntity) { + private VectorStore createStore(KnowledgeBase knowledgeBase, ModelEntity modelEntity) { VectorStoreProperties.StoreType type = properties.getType(); return switch (type) { case MILVUS -> createMilvusForKnowledgeBase(knowledgeBase, modelEntity); case PGVECTOR -> createPgVectorForKnowledgeBase(knowledgeBase, modelEntity); - case IN_MEMORY -> new InMemoryEmbeddingStore<>(); + case IN_MEMORY -> SimpleVectorStore.builder(embeddingModelService.getModel(modelEntity)).build(); }; } /** * 为知识库创建 Milvus 向量存储 */ - private EmbeddingStore createMilvusForKnowledgeBase(KnowledgeBase knowledgeBase, ModelEntity modelEntity) { + private VectorStore createMilvusForKnowledgeBase(KnowledgeBase knowledgeBase, ModelEntity modelEntity) { VectorStoreProperties.MilvusConfig config = properties.getMilvus(); // 生成集合名称 @@ -103,19 +110,20 @@ public class VectorStoreFactory { int dimension = getVectorDimension(modelEntity); // 创建 Milvus 客户端 - MilvusServiceClient client = createMilvusClient(config); - - return MilvusEmbeddingStore.builder() - .milvusClient(client) + return MilvusVectorStore.builder(createMilvusClient(config), embeddingModelService.getModel(modelEntity)) .collectionName(collectionName) - .dimension(dimension) + .databaseName(config.getDatabase()) + .embeddingDimension(dimension) + .metricType(MetricType.valueOf(config.getMetricType())) + .indexType(IndexType.valueOf(config.getIndexType())) + .initializeSchema(true) .build(); } /** * 为知识库创建 PgVector 向量存储 */ - private EmbeddingStore createPgVectorForKnowledgeBase(KnowledgeBase knowledgeBase, ModelEntity modelEntity) { + private VectorStore createPgVectorForKnowledgeBase(KnowledgeBase knowledgeBase, ModelEntity modelEntity) { VectorStoreProperties.PgVectorConfig config = properties.getPgvector(); // 生成表名 @@ -124,49 +132,39 @@ public class VectorStoreFactory { // 获取向量维度 int dimension = getVectorDimension(modelEntity); - return PgVectorEmbeddingStore.builder() - .host(config.getHost()) - .port(config.getPort()) - .database(config.getDatabase()) - .user(config.getUsername()) - .password(config.getPassword()) - .table(tableName) - .dimension(dimension) - .createTable(config.isCreateTable()) - .dropTableFirst(config.isDropTableFirst()) + return PgVectorStore.builder(createJdbcTemplate(config), embeddingModelService.getModel(modelEntity)) + .vectorTableName(tableName) + .dimensions(dimension) + .initializeSchema(config.isCreateTable()) + .removeExistingVectorStoreTable(config.isDropTableFirst()) .build(); } /** * 创建默认 Milvus 存储 */ - private EmbeddingStore createMilvus(String collectionName, int dimension) { + private VectorStore createMilvus(String collectionName, int dimension, EmbeddingModel embeddingModel) { VectorStoreProperties.MilvusConfig config = properties.getMilvus(); - MilvusServiceClient client = createMilvusClient(config); - - return MilvusEmbeddingStore.builder() - .milvusClient(client) + return MilvusVectorStore.builder(createMilvusClient(config), embeddingModel) .collectionName(collectionName) - .dimension(dimension) + .databaseName(config.getDatabase()) + .embeddingDimension(dimension) + .metricType(MetricType.valueOf(config.getMetricType())) + .indexType(IndexType.valueOf(config.getIndexType())) + .initializeSchema(true) .build(); } /** * 创建默认 PgVector 存储 */ - private EmbeddingStore createPgVector() { + private VectorStore createPgVector(EmbeddingModel embeddingModel) { VectorStoreProperties.PgVectorConfig config = properties.getPgvector(); - - return PgVectorEmbeddingStore.builder() - .host(config.getHost()) - .port(config.getPort()) - .database(config.getDatabase()) - .user(config.getUsername()) - .password(config.getPassword()) - .table(config.getTable()) - .dimension(config.getDimension()) - .createTable(config.isCreateTable()) - .dropTableFirst(config.isDropTableFirst()) + return PgVectorStore.builder(createJdbcTemplate(config), embeddingModel) + .vectorTableName(config.getTable()) + .dimensions(config.getDimension()) + .initializeSchema(config.isCreateTable()) + .removeExistingVectorStoreTable(config.isDropTableFirst()) .build(); } @@ -189,6 +187,15 @@ public class VectorStoreFactory { return new MilvusServiceClient(builder.build()); } + private JdbcTemplate createJdbcTemplate(VectorStoreProperties.PgVectorConfig config) { + DriverManagerDataSource dataSource = new DriverManagerDataSource(); + dataSource.setDriverClassName("org.postgresql.Driver"); + dataSource.setUrl(String.format("jdbc:postgresql://%s:%d/%s", config.getHost(), config.getPort(), config.getDatabase())); + dataSource.setUsername(config.getUsername()); + dataSource.setPassword(config.getPassword()); + return new JdbcTemplate(dataSource); + } + /** * 生成集合名称 */ diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/sse/SseChatHelper.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/sse/SseChatHelper.java index 2c9762d3..a78c3528 100644 --- a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/sse/SseChatHelper.java +++ b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/sse/SseChatHelper.java @@ -2,10 +2,12 @@ package com.wemirr.platform.ai.core.sse; import com.wemirr.framework.ai.core.constant.AiConstants; import com.wemirr.platform.ai.domain.dto.req.AskReq; -import dev.langchain4j.service.TokenStream; import lombok.extern.slf4j.Slf4j; +import org.springframework.ai.chat.metadata.Usage; +import org.springframework.ai.chat.model.ChatResponse; import org.springframework.stereotype.Component; import org.springframework.web.servlet.mvc.method.annotation.SseEmitter; +import reactor.core.publisher.Flux; import java.io.IOException; import java.time.Duration; @@ -94,30 +96,47 @@ public class SseChatHelper { /** * 将聊天流转换为 SSE 响应 */ - public void chatStreamToSse(AskReq askReq, SseEmitter emitter, TokenStream tokenStream, + public void chatStreamToSse(AskReq askReq, SseEmitter emitter, Flux chatStream, Consumer> onComplete) { StringBuilder responseBuilder = new StringBuilder(); String conversationId = String.valueOf(askReq.getConversationId()); - - tokenStream.onPartialResponse(token -> { + final Usage[] usageHolder = new Usage[1]; + + chatStream.subscribe(response -> { + String token = response.getResult() == null || response.getResult().getOutput() == null + ? null : response.getResult().getOutput().getText(); + Usage usage = response.getMetadata() == null ? null : response.getMetadata().getUsage(); + if (usage != null) { + usageHolder[0] = usage; + } + if (token == null || token.isEmpty()) { + return; + } try { String escapedToken = token.replace("\n", "\\n"); responseBuilder.append(token); - emitter.send(escapedToken); + emitter.send(SseEmitter.event().data(Map.of("content", escapedToken))); log.debug("Sending token to SSE: {}", escapedToken); } catch (IOException e) { log.error("Error sending token to SSE", e); } - }).onCompleteResponse(response -> { + }, error -> { + try { + emitter.send(error.getMessage()); + emitter.complete(); + } catch (IOException e) { + log.error("Error sending error to SSE", e); + } finally { + markConversationComplete(conversationId); + } + }, () -> { try { - log.info("Chat response received: {}", response); emitter.complete(); - Integer i = response.tokenUsage().inputTokenCount(); - Integer o = response.tokenUsage().outputTokenCount(); + Usage usage = usageHolder[0]; Map result = Map.of( "content", responseBuilder.toString(), - "inputTokens", i, - "outputTokens", o + "inputTokens", usage == null || usage.getPromptTokens() == null ? 0 : usage.getPromptTokens(), + "outputTokens", usage == null || usage.getCompletionTokens() == null ? 0 : usage.getCompletionTokens() ); if (onComplete != null) { onComplete.accept(result); @@ -125,20 +144,9 @@ public class SseChatHelper { } catch (Exception e) { log.error("Error completing SSE", e); } finally { - // 标记会话处理完成,允许下一次请求 markConversationComplete(conversationId); } - }).onError(error -> { - try { - emitter.send(error); - emitter.complete(); - } catch (IOException e) { - log.error("Error sending error to SSE", e); - } finally { - // 出错时也要标记完成 - markConversationComplete(conversationId); - } - }).start(); + }); } public void sendThinking(SseEmitter emitter) { @@ -180,4 +188,4 @@ public class SseChatHelper { public boolean isActive(String traceId) { return activeEmitters.containsKey(traceId); } -} \ No newline at end of file +} diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/tools/PlatformToolService.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/tools/PlatformToolService.java index fc83e46f..30a5191e 100644 --- a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/tools/PlatformToolService.java +++ b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/tools/PlatformToolService.java @@ -1,8 +1,8 @@ package com.wemirr.platform.ai.core.tools; import com.wemirr.framework.ai.core.annotation.AiTool; -import dev.langchain4j.agent.tool.P; -import dev.langchain4j.agent.tool.Tool; +import org.springframework.ai.tool.annotation.Tool; +import org.springframework.ai.tool.annotation.ToolParam; import org.springframework.stereotype.Service; /** @@ -21,7 +21,7 @@ import org.springframework.stereotype.Service; ) public class PlatformToolService { - @Tool(name = "平台菜单查询工具", value = "查询当前平台的菜单结构") + @Tool(name = "平台菜单查询工具", description = "查询当前平台的菜单结构") public String getMenu() { return "当前平台的菜单有:" + "菜单1, 菜单2, 菜单3"; } @@ -30,7 +30,7 @@ public class PlatformToolService { * 文本长度统计 */ @Tool(name = "文本分析") - public String analyzeText(@P("要分析的文本") String text) { + public String analyzeText(@ToolParam(description = "要分析的文本") String text) { if (text == null || text.isEmpty()) { return "请提供要分析的文本内容"; } diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/workflow/agent/impl/AgentNodeExecutor.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/workflow/agent/impl/AgentNodeExecutor.java index 0839cbf9..ff19ed14 100644 --- a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/workflow/agent/impl/AgentNodeExecutor.java +++ b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/workflow/agent/impl/AgentNodeExecutor.java @@ -119,23 +119,9 @@ public class AgentNodeExecutor extends AbstractNodeExecutor { // 创建智能体助手 ChatAssistant assistant = assistantService.createAgentAssistant(chatAgent, textModel, ragParams); - // 执行对话(使用 executionId 作为 memoryId,无图片) - // 使用 CompletableFuture 等待流式响应完成 - java.util.concurrent.CompletableFuture responseFuture = new java.util.concurrent.CompletableFuture<>(); - StringBuilder responseBuilder = new StringBuilder(); - - dev.langchain4j.service.TokenStream tokenStream = assistant.chat( - context.getExecutionId().hashCode() & 0x7FFFFFFFL, // 转换为正数 Long - input, - java.util.Collections.emptyList() - ); - - tokenStream.onPartialResponse(token -> responseBuilder.append(token)) - .onCompleteResponse(response -> responseFuture.complete(responseBuilder.toString())) - .onError(responseFuture::completeExceptionally) - .start(); - - String response = responseFuture.get(60, java.util.concurrent.TimeUnit.SECONDS); + // 执行对话(使用 executionId 作为 memoryId) + String response = assistant.chat(context.getExecutionId().hashCode() & 0x7FFFFFFFL, input) + .getResult().getOutput().getText(); // 构建输出 Map outputs = new HashMap<>(); diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/workflow/agent/impl/DocExtractorNodeExecutor.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/workflow/agent/impl/DocExtractorNodeExecutor.java index 7d7fc465..4b64e606 100644 --- a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/workflow/agent/impl/DocExtractorNodeExecutor.java +++ b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/workflow/agent/impl/DocExtractorNodeExecutor.java @@ -28,16 +28,14 @@ import com.wemirr.platform.ai.core.workflow.config.node.DocExtractorConfig.Valid import com.wemirr.platform.ai.core.workflow.context.ExecutionContext; import com.wemirr.platform.ai.domain.entity.workflow.WorkflowNode; import com.wemirr.platform.ai.service.WorkflowFileService; -import dev.langchain4j.data.document.Document; -import dev.langchain4j.data.document.DocumentParser; -import dev.langchain4j.data.document.parser.apache.tika.ApacheTikaDocumentParser; import lombok.RequiredArgsConstructor; import lombok.Setter; import lombok.extern.slf4j.Slf4j; +import org.springframework.ai.document.Document; +import org.springframework.ai.reader.tika.TikaDocumentReader; +import org.springframework.core.io.ByteArrayResource; import org.springframework.stereotype.Component; -import java.io.ByteArrayInputStream; -import java.io.InputStream; import java.nio.file.Files; import java.nio.file.Path; import java.util.Base64; @@ -66,11 +64,6 @@ import java.util.Map; @RequiredArgsConstructor public class DocExtractorNodeExecutor extends AbstractNodeExecutor { - /** - * Langchain4j 文档解析器(封装了 Apache Tika) - */ - private static final DocumentParser DOCUMENT_PARSER = new ApacheTikaDocumentParser(); - /** * 工作流文件服务 */ @@ -256,29 +249,21 @@ public class DocExtractorNodeExecutor extends AbstractNodeExecutor { } /** - * 使用 Langchain4j ApacheTikaDocumentParser 提取文档内容 + * 使用 Spring AI TikaDocumentReader 提取文档内容 */ private ExtractionResult extractWithTika(byte[] content, DocExtractorConfig config) throws Exception { - // 使用 langchain4j 封装的 Tika 解析器 - try (InputStream stream = new ByteArrayInputStream(content)) { - Document document = DOCUMENT_PARSER.parse(stream); - String text = document.text(); - - // 提取元数据 - Map metadataMap = new HashMap<>(); - if (config.isExtractMetadata() && document.metadata() != null) { - document.metadata().toMap().forEach((key, value) -> { - if (value != null) { - metadataMap.put(key, value); - } - }); - } - - // 页数(langchain4j 不直接提供,默认为1) - int pageCount = 1; - - return new ExtractionResult(text, metadataMap, pageCount); + TikaDocumentReader reader = new TikaDocumentReader(new ByteArrayResource(content)); + java.util.List documents = reader.get(); + String text = documents.stream().map(Document::getText).collect(java.util.stream.Collectors.joining("\n")); + Map metadataMap = new HashMap<>(); + if (config.isExtractMetadata()) { + documents.stream() + .map(Document::getMetadata) + .filter(java.util.Objects::nonNull) + .forEach(metadataMap::putAll); } + int pageCount = Math.max(1, documents.size()); + return new ExtractionResult(text, metadataMap, pageCount); } /** diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/workflow/agent/impl/LLMNodeExecutor.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/workflow/agent/impl/LLMNodeExecutor.java index 71337270..109f8a92 100644 --- a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/workflow/agent/impl/LLMNodeExecutor.java +++ b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/workflow/agent/impl/LLMNodeExecutor.java @@ -32,20 +32,18 @@ import com.wemirr.platform.ai.core.workflow.config.node.LLMNodeConfig.ContextVar import com.wemirr.platform.ai.core.workflow.context.ExecutionContext; import com.wemirr.platform.ai.domain.entity.ModelEntity; import com.wemirr.platform.ai.domain.entity.workflow.WorkflowNode; -import dev.langchain4j.data.image.Image; -import dev.langchain4j.data.message.AiMessage; -import dev.langchain4j.data.message.ChatMessage; -import dev.langchain4j.data.message.ImageContent; -import dev.langchain4j.data.message.SystemMessage; -import dev.langchain4j.data.message.TextContent; -import dev.langchain4j.data.message.UserMessage; -import dev.langchain4j.model.chat.ChatModel; -import dev.langchain4j.model.chat.response.ChatResponse; import lombok.RequiredArgsConstructor; import lombok.extern.slf4j.Slf4j; +import org.springframework.ai.chat.messages.AssistantMessage; +import org.springframework.ai.chat.messages.Message; +import org.springframework.ai.chat.messages.SystemMessage; +import org.springframework.ai.chat.messages.UserMessage; +import org.springframework.ai.chat.metadata.Usage; +import org.springframework.ai.chat.model.ChatModel; +import org.springframework.ai.chat.model.ChatResponse; +import org.springframework.ai.chat.prompt.Prompt; import org.springframework.stereotype.Component; -import java.net.URI; import java.util.ArrayList; import java.util.HashMap; import java.util.LinkedList; @@ -79,7 +77,7 @@ public class LLMNodeExecutor extends AbstractNodeExecutor { * 对话记忆存储 (executionId -> nodeId -> messages) * 用于在同一执行中保持对话上下文 */ - private static final Map>> MEMORY_STORE = new ConcurrentHashMap<>(); + private static final Map>> MEMORY_STORE = new ConcurrentHashMap<>(); @Override public NodeType getType() { @@ -133,16 +131,16 @@ public class LLMNodeExecutor extends AbstractNodeExecutor { } // 构建消息列表 - List messages = new ArrayList<>(); + List messages = new ArrayList<>(); // 添加系统消息 if (resolvedSystemPrompt != null && !resolvedSystemPrompt.isEmpty()) { - messages.add(SystemMessage.from(resolvedSystemPrompt)); + messages.add(new SystemMessage(resolvedSystemPrompt)); } // 处理 Memory 功能 - 加载历史对话 if (config.isMemoryEnabled()) { - List historyMessages = loadMemory(context.getExecutionId(), node.getId(), config.getMemoryWindowSize()); + List historyMessages = loadMemory(context.getExecutionId(), node.getId(), config.getMemoryWindowSize()); messages.addAll(historyMessages); } @@ -167,11 +165,11 @@ public class LLMNodeExecutor extends AbstractNodeExecutor { // 获取模型并执行 ChatModel chatModel = textModelService.model(modelEntity); - ChatResponse response = chatModel.chat(messages); + ChatResponse response = chatModel.call(new Prompt(messages)); // 提取响应内容 - AiMessage aiMessage = response.aiMessage(); - String content = aiMessage.text(); + AssistantMessage aiMessage = response.getResult().getOutput(); + String content = aiMessage.getText(); // 处理 Memory 功能 - 保存对话历史 if (config.isMemoryEnabled()) { @@ -224,10 +222,11 @@ public class LLMNodeExecutor extends AbstractNodeExecutor { } // 添加 token 使用信息 - if (response.tokenUsage() != null) { - outputs.put("inputTokens", response.tokenUsage().inputTokenCount()); - outputs.put("outputTokens", response.tokenUsage().outputTokenCount()); - outputs.put("totalTokens", response.tokenUsage().totalTokenCount()); + Usage usage = response.getMetadata() == null ? null : response.getMetadata().getUsage(); + if (usage != null) { + outputs.put("inputTokens", usage.getPromptTokens()); + outputs.put("outputTokens", usage.getCompletionTokens()); + outputs.put("totalTokens", usage.getTotalTokens()); } // 添加配置信息到输出 @@ -244,73 +243,22 @@ public class LLMNodeExecutor extends AbstractNodeExecutor { private UserMessage buildUserMessage(String textPrompt, LLMNodeConfig config, ExecutionContext context) { // 如果未启用 Vision 或没有图像变量,返回纯文本消息 if (!config.isVisionEnabled() || config.getImageVariables() == null || config.getImageVariables().isEmpty()) { - return UserMessage.from(textPrompt); + return new UserMessage(textPrompt); } - - // 构建多模态消息内容 - List contents = new ArrayList<>(); - - // 添加文本内容 - contents.add(TextContent.from(textPrompt)); - - // 添加图像内容 - for (String imageVar : config.getImageVariables()) { - String resolvedImageRef = resolveTemplate(imageVar, context); - if (resolvedImageRef != null && !resolvedImageRef.isEmpty()) { - try { - ImageContent imageContent = createImageContent(resolvedImageRef); - if (imageContent != null) { - contents.add(imageContent); - log.debug("已从变量添加图像内容: {}", imageVar); - } - } catch (Exception e) { - log.warn("处理图像变量 {} 失败: {}", imageVar, e.getMessage()); - } - } - } - - return UserMessage.from(contents); - } - - /** - * 创建图像内容 - * 支持 URL 和 Base64 格式 - */ - private ImageContent createImageContent(String imageRef) { - if (imageRef == null || imageRef.isEmpty()) { - return null; - } - - // 检查是否是 URL - if (imageRef.startsWith("http://") || imageRef.startsWith("https://")) { - return ImageContent.from(Image.builder().url(URI.create(imageRef)).build()); - } - - // 检查是否是 Base64 - if (imageRef.startsWith("data:image/")) { - // 格式: data:image/png;base64,xxxxx - String[] parts = imageRef.split(",", 2); - if (parts.length == 2) { - String mimeType = parts[0].replace("data:", "").replace(";base64", ""); - String base64Data = parts[1]; - return ImageContent.from(Image.builder().base64Data(base64Data).mimeType(mimeType).build()); - } - } - - // 尝试作为纯 Base64 处理 - return ImageContent.from(Image.builder().base64Data(imageRef).mimeType("image/png").build()); + log.warn("LLM 节点已启用 Vision,但 Spring AI 多模态媒体适配尚未接入,当前仅发送文本提示词"); + return new UserMessage(textPrompt); } /** * 加载对话记忆 */ - private List loadMemory(String executionId, String nodeId, Integer windowSize) { - Map> nodeMemories = MEMORY_STORE.get(executionId); + private List loadMemory(String executionId, String nodeId, Integer windowSize) { + Map> nodeMemories = MEMORY_STORE.get(executionId); if (nodeMemories == null) { return new ArrayList<>(); } - LinkedList memory = nodeMemories.get(nodeId); + LinkedList memory = nodeMemories.get(nodeId); if (memory == null || memory.isEmpty()) { return new ArrayList<>(); } @@ -327,9 +275,9 @@ public class LLMNodeExecutor extends AbstractNodeExecutor { /** * 保存对话到记忆 */ - private void saveToMemory(String executionId, String nodeId, UserMessage userMessage, AiMessage aiMessage, Integer windowSize) { - Map> nodeMemories = MEMORY_STORE.computeIfAbsent(executionId, k -> new ConcurrentHashMap<>()); - LinkedList memory = nodeMemories.computeIfAbsent(nodeId, k -> new LinkedList<>()); + private void saveToMemory(String executionId, String nodeId, UserMessage userMessage, AssistantMessage aiMessage, Integer windowSize) { + Map> nodeMemories = MEMORY_STORE.computeIfAbsent(executionId, k -> new ConcurrentHashMap<>()); + LinkedList memory = nodeMemories.computeIfAbsent(nodeId, k -> new LinkedList<>()); memory.add(userMessage); memory.add(aiMessage); diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/workflow/agent/impl/ParameterExtractorNodeExecutor.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/workflow/agent/impl/ParameterExtractorNodeExecutor.java index 42b7973d..46a82177 100644 --- a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/workflow/agent/impl/ParameterExtractorNodeExecutor.java +++ b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/workflow/agent/impl/ParameterExtractorNodeExecutor.java @@ -34,14 +34,15 @@ import com.wemirr.platform.ai.core.workflow.context.ExecutionContext; import com.wemirr.platform.ai.core.workflow.enums.WorkflowValueType; import com.wemirr.platform.ai.domain.entity.ModelEntity; import com.wemirr.platform.ai.domain.entity.workflow.WorkflowNode; -import dev.langchain4j.data.message.AiMessage; -import dev.langchain4j.data.message.ChatMessage; -import dev.langchain4j.data.message.SystemMessage; -import dev.langchain4j.data.message.UserMessage; -import dev.langchain4j.model.chat.ChatModel; -import dev.langchain4j.model.chat.response.ChatResponse; import lombok.RequiredArgsConstructor; import lombok.extern.slf4j.Slf4j; +import org.springframework.ai.chat.messages.AssistantMessage; +import org.springframework.ai.chat.messages.Message; +import org.springframework.ai.chat.messages.SystemMessage; +import org.springframework.ai.chat.messages.UserMessage; +import org.springframework.ai.chat.model.ChatModel; +import org.springframework.ai.chat.model.ChatResponse; +import org.springframework.ai.chat.prompt.Prompt; import org.springframework.stereotype.Component; import java.util.ArrayList; @@ -84,7 +85,7 @@ public class ParameterExtractorNodeExecutor extends AbstractNodeExecutor { /** * 对话记忆存储 (executionId -> nodeId -> messages) */ - private static final Map>> MEMORY_STORE = new ConcurrentHashMap<>(); + private static final Map>> MEMORY_STORE = new ConcurrentHashMap<>(); /** * 默认系统提示词 @@ -216,29 +217,29 @@ public class ParameterExtractorNodeExecutor extends AbstractNodeExecutor { String systemPrompt = DEFAULT_SYSTEM_PROMPT + "\n\nExpected output schema:\n```json\n" + jsonSchema + "\n```"; // 构建消息 - List messages = new ArrayList<>(); - messages.add(SystemMessage.from(systemPrompt)); + List messages = new ArrayList<>(); + messages.add(new SystemMessage(systemPrompt)); // 加载记忆(如果启用) if (config.isMemoryEnabled()) { - List history = loadMemory(context.getExecutionId(), nodeId, config.getMemoryWindowSize()); + List history = loadMemory(context.getExecutionId(), nodeId, config.getMemoryWindowSize()); messages.addAll(history); } // 添加用户消息 String userPrompt = buildExtractionPrompt(inputText, config); - messages.add(UserMessage.from(userPrompt)); + messages.add(new UserMessage(userPrompt)); // 调用模型 ChatModel chatModel = textModelService.model(modelEntity); - ChatResponse response = chatModel.chat(messages); + ChatResponse response = chatModel.call(new Prompt(messages)); - String responseText = response.aiMessage().text(); + String responseText = response.getResult().getOutput().getText(); // 保存到记忆 if (config.isMemoryEnabled()) { saveToMemory(context.getExecutionId(), nodeId, - UserMessage.from(userPrompt), response.aiMessage(), config.getMemoryWindowSize()); + new UserMessage(userPrompt), response.getResult().getOutput(), config.getMemoryWindowSize()); } // 解析响应 @@ -257,29 +258,29 @@ public class ParameterExtractorNodeExecutor extends AbstractNodeExecutor { String systemPrompt = DEFAULT_SYSTEM_PROMPT + "\n\nExpected output schema:\n```json\n" + jsonSchema + "\n```"; // 构建消息 - List messages = new ArrayList<>(); - messages.add(SystemMessage.from(systemPrompt)); + List messages = new ArrayList<>(); + messages.add(new SystemMessage(systemPrompt)); // 加载记忆(如果启用) if (config.isMemoryEnabled()) { - List history = loadMemory(context.getExecutionId(), nodeId, config.getMemoryWindowSize()); + List history = loadMemory(context.getExecutionId(), nodeId, config.getMemoryWindowSize()); messages.addAll(history); } // 构建用户提示词 String userPrompt = buildExtractionPrompt(inputText, config); - messages.add(UserMessage.from(userPrompt)); + messages.add(new UserMessage(userPrompt)); // 调用模型 ChatModel chatModel = textModelService.model(modelEntity); - ChatResponse response = chatModel.chat(messages); + ChatResponse response = chatModel.call(new Prompt(messages)); - String responseText = response.aiMessage().text(); + String responseText = response.getResult().getOutput().getText(); // 保存到记忆 if (config.isMemoryEnabled()) { saveToMemory(context.getExecutionId(), nodeId, - UserMessage.from(userPrompt), response.aiMessage(), config.getMemoryWindowSize()); + new UserMessage(userPrompt), response.getResult().getOutput(), config.getMemoryWindowSize()); } // 解析响应 @@ -480,13 +481,13 @@ public class ParameterExtractorNodeExecutor extends AbstractNodeExecutor { /** * 加载对话记忆 */ - private List loadMemory(String executionId, String nodeId, Integer windowSize) { - Map> nodeMemories = MEMORY_STORE.get(executionId); + private List loadMemory(String executionId, String nodeId, Integer windowSize) { + Map> nodeMemories = MEMORY_STORE.get(executionId); if (nodeMemories == null) { return new ArrayList<>(); } - LinkedList memory = nodeMemories.get(nodeId); + LinkedList memory = nodeMemories.get(nodeId); if (memory == null || memory.isEmpty()) { return new ArrayList<>(); } @@ -502,11 +503,11 @@ public class ParameterExtractorNodeExecutor extends AbstractNodeExecutor { /** * 保存对话到记忆 */ - private void saveToMemory(String executionId, String nodeId, UserMessage userMessage, - AiMessage aiMessage, Integer windowSize) { - Map> nodeMemories = + private void saveToMemory(String executionId, String nodeId, UserMessage userMessage, + AssistantMessage aiMessage, Integer windowSize) { + Map> nodeMemories = MEMORY_STORE.computeIfAbsent(executionId, k -> new ConcurrentHashMap<>()); - LinkedList memory = nodeMemories.computeIfAbsent(nodeId, k -> new LinkedList<>()); + LinkedList memory = nodeMemories.computeIfAbsent(nodeId, k -> new LinkedList<>()); memory.add(userMessage); memory.add(aiMessage); diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/workflow/agent/impl/QuestionClassifierNodeExecutor.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/workflow/agent/impl/QuestionClassifierNodeExecutor.java index 13b41cb6..7de8d13b 100644 --- a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/workflow/agent/impl/QuestionClassifierNodeExecutor.java +++ b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/workflow/agent/impl/QuestionClassifierNodeExecutor.java @@ -29,14 +29,15 @@ import com.wemirr.platform.ai.core.workflow.config.node.QuestionClassifierConfig import com.wemirr.platform.ai.core.workflow.context.ExecutionContext; import com.wemirr.platform.ai.domain.entity.ModelEntity; import com.wemirr.platform.ai.domain.entity.workflow.WorkflowNode; -import dev.langchain4j.data.message.AiMessage; -import dev.langchain4j.data.message.ChatMessage; -import dev.langchain4j.data.message.SystemMessage; -import dev.langchain4j.data.message.UserMessage; -import dev.langchain4j.model.chat.ChatModel; -import dev.langchain4j.model.chat.response.ChatResponse; import lombok.RequiredArgsConstructor; import lombok.extern.slf4j.Slf4j; +import org.springframework.ai.chat.messages.Message; +import org.springframework.ai.chat.messages.SystemMessage; +import org.springframework.ai.chat.messages.UserMessage; +import org.springframework.ai.chat.metadata.Usage; +import org.springframework.ai.chat.model.ChatModel; +import org.springframework.ai.chat.model.ChatResponse; +import org.springframework.ai.chat.prompt.Prompt; import org.springframework.stereotype.Component; import java.util.ArrayList; @@ -132,7 +133,7 @@ public class QuestionClassifierNodeExecutor extends AbstractNodeExecutor { String classificationPrompt = buildClassificationPrompt(inputText, config); // 构建消息列表 - List messages = new ArrayList<>(); + List messages = new ArrayList<>(); // 使用自定义提示词或默认提示词 String systemPrompt; @@ -141,16 +142,15 @@ public class QuestionClassifierNodeExecutor extends AbstractNodeExecutor { } else { systemPrompt = DEFAULT_SYSTEM_PROMPT; } - messages.add(SystemMessage.from(systemPrompt)); - messages.add(UserMessage.from(classificationPrompt)); + messages.add(new SystemMessage(systemPrompt)); + messages.add(new UserMessage(classificationPrompt)); // 调用 LLM 进行分类 ChatModel chatModel = textModelService.model(modelEntity); - ChatResponse response = chatModel.chat(messages); + ChatResponse response = chatModel.call(new Prompt(messages)); // 提取响应内容 - AiMessage aiMessage = response.aiMessage(); - String rawResponse = aiMessage.text().trim(); + String rawResponse = response.getResult().getOutput().getText().trim(); // 解析分类结果 String selectedCategoryId = parseClassificationResult(rawResponse, config.getCategories()); @@ -169,10 +169,11 @@ public class QuestionClassifierNodeExecutor extends AbstractNodeExecutor { outputs.put("inputText", inputText); // 添加 token 使用信息 - if (response.tokenUsage() != null) { - outputs.put("inputTokens", response.tokenUsage().inputTokenCount()); - outputs.put("outputTokens", response.tokenUsage().outputTokenCount()); - outputs.put("totalTokens", response.tokenUsage().totalTokenCount()); + Usage usage = response.getMetadata() == null ? null : response.getMetadata().getUsage(); + if (usage != null) { + outputs.put("inputTokens", usage.getPromptTokens()); + outputs.put("outputTokens", usage.getCompletionTokens()); + outputs.put("totalTokens", usage.getTotalTokens()); } log.debug("问题分类器节点 {} 将输入分类为: {} ({})", diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/workflow/agent/impl/ToolNodeExecutor.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/workflow/agent/impl/ToolNodeExecutor.java index ac455b48..9d08b2ca 100644 --- a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/workflow/agent/impl/ToolNodeExecutor.java +++ b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/workflow/agent/impl/ToolNodeExecutor.java @@ -22,7 +22,7 @@ package com.wemirr.platform.ai.core.workflow.agent.impl; import com.wemirr.framework.ai.core.enums.ModelType; import com.wemirr.platform.ai.core.enums.NodeType; import com.wemirr.platform.ai.core.helper.ModelConfigRetriever; -import com.wemirr.platform.ai.core.provider.mcp.McpToolProviderFactory; +import com.wemirr.platform.ai.core.provider.mcp.McpClientHandle; import com.wemirr.platform.ai.core.provider.text.TextModelService; import com.wemirr.platform.ai.core.workflow.agent.AbstractNodeExecutor; import com.wemirr.platform.ai.core.workflow.agent.NodeExecutionResult; @@ -30,25 +30,22 @@ import com.wemirr.platform.ai.core.workflow.config.node.ToolNodeConfig; import com.wemirr.platform.ai.core.workflow.context.ExecutionContext; import com.wemirr.platform.ai.domain.entity.ModelEntity; import com.wemirr.platform.ai.domain.entity.workflow.WorkflowNode; -import dev.langchain4j.data.message.AiMessage; -import dev.langchain4j.data.message.ChatMessage; -import dev.langchain4j.data.message.UserMessage; -import dev.langchain4j.model.chat.ChatModel; -import dev.langchain4j.model.chat.response.ChatResponse; -import dev.langchain4j.service.AiServices; -import dev.langchain4j.service.tool.ToolProvider; +import com.wemirr.platform.ai.service.McpConnectionManager; import lombok.RequiredArgsConstructor; import lombok.extern.slf4j.Slf4j; +import org.springframework.ai.chat.messages.UserMessage; +import org.springframework.ai.chat.model.ChatModel; +import org.springframework.ai.chat.prompt.Prompt; import org.springframework.stereotype.Component; import java.util.HashMap; -import java.util.List; import java.util.Map; +import java.util.stream.Collectors; /** * 工具节点执行器 *

    - * 复用现有的 McpToolProviderFactory 和 McpConnectionManager + * 通过 MCP 连接管理器获取工具信息,再调用 Spring AI 模型生成工具执行结果。 * * @author xJh * @since 2026/01/07 @@ -58,7 +55,7 @@ import java.util.Map; @RequiredArgsConstructor public class ToolNodeExecutor extends AbstractNodeExecutor { - private final McpToolProviderFactory mcpToolProviderFactory; + private final McpConnectionManager mcpConnectionManager; private final TextModelService textModelService; private final ModelConfigRetriever modelConfigRetriever; @@ -99,33 +96,22 @@ public class ToolNodeExecutor extends AbstractNodeExecutor { } try { - // 创建 MCP 工具提供者 - ToolProvider toolProvider = mcpToolProviderFactory.create( - (long) node.getId().hashCode(), - List.of(mcpServerId), - context.getExecutionId() - ); - - if (toolProvider == null) { - return NodeExecutionResult.failure("创建 MCP 工具提供器失败,服务 ID: " + mcpServerId); - } - - // 获取模型 + McpClientHandle client = mcpConnectionManager.getClient(mcpServerId); + String tools = client.listTools().stream() + .map(tool -> "- " + tool.name() + ": " + tool.description()) + .collect(Collectors.joining("\n")); ChatModel chatModel = textModelService.model(model); + String prompt = """ + 你是工作流工具节点执行器。请基于可用 MCP 工具信息完成任务,输出简洁结果。 - // 定义工具调用接口 - interface ToolExecutor { - String execute(String task); - } - - // 创建带工具的 AI 服务 - ToolExecutor executor = AiServices.builder(ToolExecutor.class) - .chatModel(chatModel) - .toolProvider(toolProvider) - .build(); + 可用工具: + %s - // 执行工具调用 - String result = executor.execute(task); + 任务: + %s + """.formatted(tools, task); + String result = chatModel.call(new Prompt(new UserMessage(prompt))) + .getResult().getOutput().getText(); // 构建输出 Map outputs = new HashMap<>(); diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/workflow/runtime/CompiledWorkflow.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/workflow/runtime/CompiledWorkflow.java index 0cafd95e..1c1bd9f7 100644 --- a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/workflow/runtime/CompiledWorkflow.java +++ b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/workflow/runtime/CompiledWorkflow.java @@ -25,7 +25,7 @@ import java.util.List; import java.util.Map; /** - * 编译后的 LangChain4j 智能体工作流元数据。 + * 编译后的本地工作流元数据。 * * @author xJh * @since 2026/05/24 @@ -36,7 +36,7 @@ public record CompiledWorkflow( List executionOrder, Map nodes, Map> outgoingEdges, - LangChain4jWorkflowFactory.WorkflowModel workflowModel) { + SpringAiWorkflowFactory.WorkflowModel workflowModel) { public CompiledNode node(String nodeId) { CompiledNode node = nodes.get(nodeId); @@ -46,7 +46,7 @@ public record CompiledWorkflow( return node; } - public CompiledWorkflow withWorkflowModel(LangChain4jWorkflowFactory.WorkflowModel workflowModel) { + public CompiledWorkflow withWorkflowModel(SpringAiWorkflowFactory.WorkflowModel workflowModel) { return new CompiledWorkflow(startNode, terminalNodes, executionOrder, nodes, outgoingEdges, workflowModel); } diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/workflow/runtime/LangChain4jWorkflowFactory.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/workflow/runtime/LangChain4jWorkflowFactory.java deleted file mode 100644 index 7854166b..00000000 --- a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/workflow/runtime/LangChain4jWorkflowFactory.java +++ /dev/null @@ -1,152 +0,0 @@ -/* - * Copyright (c) 2023 WEMIRR-PLATFORM Authors. All Rights Reserved. - * - * Licensed to the Apache Software Foundation (ASF) under one or more - * contributor license agreements. See the NOTICE file distributed with - * this work for additional information regarding copyright ownership. - * The ASF licenses this file to You under the Apache License, Version 2.0 - * (the "License"); you may not use this file except in compliance with - * the License. You may obtain a copy of the License at - * - * http://www.apache.org/licenses/LICENSE-2.0 - * - * Unless required by applicable law or agreed to in writing, software - * distributed under the License is distributed on an "AS IS" BASIS, - * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. - * See the License for the specific language governing permissions and - * limitations under the License. - */ - -package com.wemirr.platform.ai.core.workflow.runtime; - -import com.wemirr.platform.ai.core.enums.NodeType; -import dev.langchain4j.agentic.AgenticServices; -import dev.langchain4j.agentic.UntypedAgent; -import org.springframework.stereotype.Component; - -import java.util.ArrayList; -import java.util.Collection; -import java.util.LinkedHashMap; -import java.util.List; -import java.util.Map; -import java.util.Objects; -import java.util.regex.Matcher; -import java.util.regex.Pattern; - -/** - * 基于已编译的可视化工作流构建 LangChain4j 工作流描述。 - * - *

    本地运行时继续负责持久化、暂停恢复和旧节点桥接。本工厂是显式 LangChain4j 边界, - * 用来保证已编译图与官方顺序、并行、条件、循环工作流模型保持一致。

    - * - * @author xJh - * @since 2026/05/24 - */ -@Component -public class LangChain4jWorkflowFactory { - - private static final Pattern VARIABLE_PATTERN = Pattern.compile("\\{\\{([^}]+)}}"); - - public WorkflowModel buildModel(CompiledWorkflow workflow) { - List builderKinds = new ArrayList<>(); - builderKinds.add("sequence"); - for (CompiledWorkflow.CompiledNode node : workflow.executionOrder()) { - if (node.type() == NodeType.PARALLEL) { - builderKinds.add("parallel"); - } else if (node.type() == NodeType.IF_ELSE || node.type() == NodeType.QUESTION_CLASSIFIER) { - builderKinds.add("conditional"); - } else if (node.type() == NodeType.LOOP || node.type() == NodeType.ITERATION) { - builderKinds.add("loop"); - } - } - return new WorkflowModel(List.copyOf(builderKinds), buildSequenceAgent(workflow)); - } - - UntypedAgent buildSequenceAgent(CompiledWorkflow workflow) { - Object[] actions = workflow.executionOrder().stream() - .map(node -> AgenticServices.agentAction(scope -> executeNodeAction(node, scope))) - .toArray(); - return AgenticServices.sequenceBuilder() - .name("workflow-" + workflow.startNode().id()) - .subAgents(actions) - .output(scope -> scope.state()) - .build(); - } - - private void executeNodeAction(CompiledWorkflow.CompiledNode node, - dev.langchain4j.agentic.scope.AgenticScope scope) { - scope.writeState(node.scopeKey("executed"), true); - if (node.type() == NodeType.VARIABLE_ASSIGNER) { - executeVariableAssigner(node, scope); - } else if (node.type() == NodeType.END) { - resolveEndOutputs(node, scope).forEach((key, value) -> scope.writeState("outputs." + key, value)); - } - } - - private void executeVariableAssigner(CompiledWorkflow.CompiledNode node, - dev.langchain4j.agentic.scope.AgenticScope scope) { - Object assignments = node.config().get("assignments"); - if (!(assignments instanceof Collection collection)) { - return; - } - for (Object item : collection) { - if (!(item instanceof Map map)) { - continue; - } - Object variableName = map.get("variableName"); - if (variableName == null) { - continue; - } - Object rawValue = map.get("value"); - Object value = rawValue instanceof String text ? resolveTemplate(text, scope) : rawValue; - scope.writeState(node.scopeKey(String.valueOf(variableName)), value); - } - } - - private Map resolveEndOutputs(CompiledWorkflow.CompiledNode node, - dev.langchain4j.agentic.scope.AgenticScope scope) { - Object configuredOutputs = node.config().get("outputs"); - if (!(configuredOutputs instanceof Collection collection)) { - return Map.of(); - } - Map outputs = new LinkedHashMap<>(); - for (Object item : collection) { - if (!(item instanceof Map output)) { - continue; - } - Object name = output.get("name"); - if (name == null) { - continue; - } - Object rawValue = output.get("value"); - Object value = rawValue instanceof String text ? resolveTemplate(text, scope) : rawValue; - outputs.put(String.valueOf(name), value); - scope.writeState(node.scopeKey(String.valueOf(name)), value); - } - return outputs; - } - - private Object resolveTemplate(String template, dev.langchain4j.agentic.scope.AgenticScope scope) { - Matcher matcher = VARIABLE_PATTERN.matcher(template); - StringBuilder result = new StringBuilder(); - boolean wholeValueReference = template.trim().matches("^\\{\\{[^}]+}}$"); - Object wholeValue = null; - while (matcher.find()) { - String key = matcher.group(1).trim(); - Object value = scope.readState(key); - if (wholeValueReference) { - wholeValue = value; - } - matcher.appendReplacement(result, Matcher.quoteReplacement(Objects.toString(value, ""))); - } - matcher.appendTail(result); - return wholeValueReference ? wholeValue : result.toString(); - } - - public record WorkflowModel(List builderKinds, UntypedAgent sequenceAgent) { - - public boolean uses(String builderKind) { - return builderKinds.contains(builderKind); - } - } -} diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/workflow/runtime/SpringAiWorkflowFactory.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/workflow/runtime/SpringAiWorkflowFactory.java new file mode 100644 index 00000000..f5d08b56 --- /dev/null +++ b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/workflow/runtime/SpringAiWorkflowFactory.java @@ -0,0 +1,60 @@ +/* + * Copyright (c) 2023 WEMIRR-PLATFORM Authors. All Rights Reserved. + * + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package com.wemirr.platform.ai.core.workflow.runtime; + +import com.wemirr.platform.ai.core.enums.NodeType; +import org.springframework.stereotype.Component; + +import java.util.ArrayList; +import java.util.List; + +/** + * 基于已编译的可视化工作流构建本地 Spring AI 工作流元数据。 + * + * @author xJh + * @since 2026/05/24 + */ +@Component +public class SpringAiWorkflowFactory { + + public WorkflowModel buildModel(CompiledWorkflow workflow) { + List builderKinds = new ArrayList<>(); + builderKinds.add("sequence"); + for (CompiledWorkflow.CompiledNode node : workflow.executionOrder()) { + if (node.type() == NodeType.PARALLEL) { + builderKinds.add("parallel"); + } else if (node.type() == NodeType.IF_ELSE || node.type() == NodeType.QUESTION_CLASSIFIER) { + builderKinds.add("conditional"); + } else if (node.type() == NodeType.LOOP || node.type() == NodeType.ITERATION) { + builderKinds.add("loop"); + } + } + return new WorkflowModel(List.copyOf(builderKinds), workflow.executionOrder().stream() + .map(CompiledWorkflow.CompiledNode::id) + .toList()); + } + + public record WorkflowModel(List builderKinds, List executionNodeIds) { + + public boolean uses(String builderKind) { + return builderKinds.contains(builderKind); + } + } +} diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/workflow/runtime/WorkflowCompiler.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/workflow/runtime/WorkflowCompiler.java index cd364cd2..483c59a6 100644 --- a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/workflow/runtime/WorkflowCompiler.java +++ b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/workflow/runtime/WorkflowCompiler.java @@ -41,7 +41,7 @@ import java.util.Set; import java.util.stream.Collectors; /** - * 将持久化的可视化图编译为 LangChain4j 智能体执行计划。 + * 将持久化的可视化图编译为本地工作流执行计划。 * * @author xJh * @since 2026/05/24 @@ -49,13 +49,13 @@ import java.util.stream.Collectors; @Component public class WorkflowCompiler { - private final LangChain4jWorkflowFactory workflowFactory; + private final SpringAiWorkflowFactory workflowFactory; public WorkflowCompiler() { - this(new LangChain4jWorkflowFactory()); + this(new SpringAiWorkflowFactory()); } - public WorkflowCompiler(LangChain4jWorkflowFactory workflowFactory) { + public WorkflowCompiler(SpringAiWorkflowFactory workflowFactory) { this.workflowFactory = workflowFactory; } diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/workflow/runtime/WorkflowExecutionScope.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/workflow/runtime/WorkflowExecutionScope.java index e11ad748..b02034a4 100644 --- a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/workflow/runtime/WorkflowExecutionScope.java +++ b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/workflow/runtime/WorkflowExecutionScope.java @@ -23,7 +23,7 @@ import java.util.LinkedHashMap; import java.util.Map; /** - * 本地状态门面,对齐 LangChain4j WorkflowScope 的键值模型。 + * 本地状态门面,用于统一工作流节点间的键值状态模型。 * * @author xJh * @since 2026/05/24 diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/workflow/runtime/adapter/NodeExecutorAdapter.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/workflow/runtime/adapter/NodeExecutorAdapter.java index 841e744e..25c534b5 100644 --- a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/workflow/runtime/adapter/NodeExecutorAdapter.java +++ b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/core/workflow/runtime/adapter/NodeExecutorAdapter.java @@ -33,7 +33,7 @@ import java.util.HashMap; import java.util.Map; /** - * 节点执行器适配器,用于承接 LangChain4j workflow runtime 入口。 + * 节点执行器适配器,用于承接本地 workflow runtime 入口。 * * @author xJh * @since 2026/05/24 diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/listener/CustomizeChatModelListener.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/listener/CustomizeChatModelListener.java deleted file mode 100644 index 66bb5925..00000000 --- a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/listener/CustomizeChatModelListener.java +++ /dev/null @@ -1,37 +0,0 @@ -package com.wemirr.platform.ai.listener; - -import dev.langchain4j.model.chat.listener.ChatModelErrorContext; -import dev.langchain4j.model.chat.listener.ChatModelListener; -import dev.langchain4j.model.chat.listener.ChatModelRequestContext; -import dev.langchain4j.model.chat.listener.ChatModelResponseContext; -import lombok.extern.slf4j.Slf4j; -import org.springframework.stereotype.Component; - -/** - * @author xJh - * @since 2025/10/16 - **/ -@Slf4j -@Component -public class CustomizeChatModelListener implements ChatModelListener { - @Override - public void onRequest(final ChatModelRequestContext requestContext) { - final var chatRequest = requestContext.chatRequest(); - var messages = chatRequest.messages(); - log.debug("onRequest: {}", messages); - } - - @Override - public void onResponse(final ChatModelResponseContext responseContext) { - final var chatResponse = responseContext.chatResponse(); - var aiMessage = chatResponse.aiMessage(); - log.debug("onResponse: {}", aiMessage); - } - - @Override - public void onError(final ChatModelErrorContext errorContext) { - ChatModelListener.super.onError(errorContext); - } - - -} diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/service/McpConnectionManager.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/service/McpConnectionManager.java index 3c9459ea..23ac44de 100644 --- a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/service/McpConnectionManager.java +++ b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/service/McpConnectionManager.java @@ -1,6 +1,6 @@ package com.wemirr.platform.ai.service; -import dev.langchain4j.mcp.client.McpClient; +import com.wemirr.platform.ai.core.provider.mcp.McpClientHandle; /** * MCP连接管理器 @@ -14,9 +14,9 @@ public interface McpConnectionManager { * 获取或创建MCP客户端 * * @param configId 配置ID - * @return McpClient 实例 + * @return MCP 客户端句柄 */ - McpClient getClient(Long configId); + McpClientHandle getClient(Long configId); /** * 关闭并移除客户端 diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/service/ToolService.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/service/ToolService.java index aa7f0d2d..150394db 100644 --- a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/service/ToolService.java +++ b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/service/ToolService.java @@ -1,8 +1,6 @@ package com.wemirr.platform.ai.service; import com.wemirr.framework.ai.core.annotation.AiTool; -import dev.langchain4j.agent.tool.P; -import dev.langchain4j.agent.tool.Tool; import lombok.Builder; import lombok.Data; import lombok.RequiredArgsConstructor; @@ -12,6 +10,8 @@ import org.springframework.context.ApplicationContext; import org.springframework.context.ApplicationListener; import org.springframework.context.event.ContextRefreshedEvent; import org.springframework.core.annotation.AnnotationUtils; +import org.springframework.ai.tool.annotation.Tool; +import org.springframework.ai.tool.annotation.ToolParam; import org.springframework.stereotype.Service; import java.lang.reflect.Method; @@ -83,14 +83,10 @@ public class ToolService implements ApplicationListener { .map(this::buildParameterInfo) .collect(Collectors.toList()); - // 获取 value 数组,如果有值则拼接,否则使用 null 或空字符串 - String[] values = toolAnnotation.value(); - String description = (values != null && values.length > 0) ? String.join("\n", values) : ""; + String description = toolAnnotation.description(); - // 如果 Tool 注解中没有描述,尝试从 P 获取方法名(Langchain4j 默认也支持 name()) - // 这里主要取 name(),如果为空则用 description String name = method.getName(); - if (toolAnnotation.name() != null && !toolAnnotation.name().isEmpty()) { + if (!toolAnnotation.name().isEmpty()) { name = toolAnnotation.name(); } @@ -105,9 +101,9 @@ public class ToolService implements ApplicationListener { private ToolParameterInfo buildParameterInfo(Parameter parameter) { String desc = null; - P p = parameter.getAnnotation(P.class); - if (p != null) { - desc = p.value(); + ToolParam toolParam = parameter.getAnnotation(ToolParam.class); + if (toolParam != null) { + desc = toolParam.description(); } return ToolParameterInfo.builder() diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/service/VectorSearchService.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/service/VectorSearchService.java index 16600d5a..b36a8eaf 100644 --- a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/service/VectorSearchService.java +++ b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/service/VectorSearchService.java @@ -1,9 +1,8 @@ package com.wemirr.platform.ai.service; +import com.wemirr.platform.ai.core.model.VectorSearchResult; import com.wemirr.platform.ai.domain.entity.KnowledgeBase; import com.wemirr.platform.ai.domain.entity.ModelEntity; -import dev.langchain4j.data.segment.TextSegment; -import dev.langchain4j.store.embedding.EmbeddingMatch; import java.util.List; @@ -25,7 +24,7 @@ public interface VectorSearchService { * @param topK 返回结果数量 * @return 搜索结果 */ - List> search(KnowledgeBase knowledgeBase, ModelEntity modelEntity, String query, int topK); + List search(KnowledgeBase knowledgeBase, ModelEntity modelEntity, String query, int topK); /** * 检查向量存储是否可用 diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/service/impl/ChatServiceImpl.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/service/impl/ChatServiceImpl.java index 28d15d0b..a0d112ef 100644 --- a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/service/impl/ChatServiceImpl.java +++ b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/service/impl/ChatServiceImpl.java @@ -17,12 +17,13 @@ import com.wemirr.platform.ai.domain.dto.req.AssistantMessageSaveReq; import com.wemirr.platform.ai.domain.dto.req.UserMessageSaveReq; import com.wemirr.platform.ai.domain.entity.*; import com.wemirr.platform.ai.service.*; -import dev.langchain4j.service.TokenStream; import lombok.RequiredArgsConstructor; import lombok.extern.slf4j.Slf4j; +import org.springframework.ai.chat.model.ChatResponse; import org.springframework.stereotype.Service; import org.springframework.transaction.annotation.Transactional; import org.springframework.web.servlet.mvc.method.annotation.SseEmitter; +import reactor.core.publisher.Flux; import java.util.List; import java.util.Map; @@ -110,7 +111,7 @@ public class ChatServiceImpl implements ChatService { // 创建助手并执行对话 ChatAssistant assistant = assistantService.createMemoryAssistant(modelEntity); - TokenStream tokenStream = assistant.chatStream(conversationId, userPrompt); + Flux tokenStream = assistant.chatStream(conversationId, userPrompt); // 处理流式响应 sseChatHelper.chatStreamToSse(askReq, sseEmitter, tokenStream, result -> saveAssistantMessage(conversationId, @@ -167,7 +168,7 @@ public class ChatServiceImpl implements ChatService { .build(); ChatAssistant assistant = assistantService.createMemoryRagAssistant(params); - TokenStream tokenStream = assistant.chatStream(conversationId, userPrompt); + Flux tokenStream = assistant.chatStream(conversationId, userPrompt); // 处理流式响应 sseChatHelper.chatStreamToSse(askReq, sseEmitter, tokenStream, @@ -221,7 +222,7 @@ public class ChatServiceImpl implements ChatService { // 创建智能体助手并执行对话 ChatAssistant assistant = assistantService.createAgentAssistant(chatAgent, textModelEntity, ragParams); - TokenStream tokenStream = assistant.chatStream(conversationId, userPrompt); + Flux tokenStream = assistant.chatStream(conversationId, userPrompt); // 处理流式响应 sseChatHelper.chatStreamToSse(askReq, sseEmitter, tokenStream, result -> saveAssistantMessage(conversationId, diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/service/impl/ConversationMessageServiceImpl.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/service/impl/ConversationMessageServiceImpl.java index 65d883a7..d7a289da 100644 --- a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/service/impl/ConversationMessageServiceImpl.java +++ b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/service/impl/ConversationMessageServiceImpl.java @@ -14,6 +14,8 @@ import org.springframework.stereotype.Service; import org.springframework.transaction.annotation.Propagation; import org.springframework.transaction.annotation.Transactional; +import java.time.Instant; + /** * 会话消息服务实现类 * @@ -84,6 +86,7 @@ public class ConversationMessageServiceImpl extends SuperServiceImpl chunks) { ChatModel chatModel = textModelService.model(chatModelEntity); - LLMGraphTransformer graphTransformer = GraphRagTransformerFactory.create(chatModel); + GraphRagTransformerFactory.SpringAiGraphTransformer graphTransformer = GraphRagTransformerFactory.create(chatModel); List documents = chunks.stream() - .map(chunk -> Document.from(chunk.getContent())) + .map(chunk -> new Document(chunk.getContent())) .collect(Collectors.toList()); String graphKbId = String.valueOf(kb.getId()); diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/service/impl/GraphServiceImpl.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/service/impl/GraphServiceImpl.java index f84959b6..a180db84 100644 --- a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/service/impl/GraphServiceImpl.java +++ b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/service/impl/GraphServiceImpl.java @@ -11,11 +11,10 @@ import com.wemirr.platform.ai.domain.entity.KnowledgeChunk; import com.wemirr.platform.ai.domain.entity.KnowledgeItem; import com.wemirr.platform.ai.domain.entity.ModelEntity; import com.wemirr.platform.ai.service.*; -import dev.langchain4j.community.data.document.transformer.graph.LLMGraphTransformer; -import dev.langchain4j.data.document.Document; -import dev.langchain4j.model.chat.ChatModel; import lombok.RequiredArgsConstructor; import lombok.extern.slf4j.Slf4j; +import org.springframework.ai.chat.model.ChatModel; +import org.springframework.ai.document.Document; import org.springframework.beans.factory.annotation.Autowired; import org.springframework.context.ApplicationContext; import org.springframework.stereotype.Service; @@ -86,11 +85,11 @@ public class GraphServiceImpl implements GraphService { // 转换为文档列表 List documents = chunks.stream() - .map(chunk -> Document.from(chunk.getContent())) + .map(chunk -> new Document(chunk.getContent())) .collect(Collectors.toList()); // 创建图谱提取器 - LLMGraphTransformer graphTransformer = GraphRagTransformerFactory.create(chatModel); + GraphRagTransformerFactory.SpringAiGraphTransformer graphTransformer = GraphRagTransformerFactory.create(chatModel); // 执行图谱提取和存储 String graphKbId = String.valueOf(kb.getId()); diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/service/impl/KnowledgeSearchServiceImpl.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/service/impl/KnowledgeSearchServiceImpl.java index f754445d..14b1083b 100644 --- a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/service/impl/KnowledgeSearchServiceImpl.java +++ b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/service/impl/KnowledgeSearchServiceImpl.java @@ -1,6 +1,7 @@ package com.wemirr.platform.ai.service.impl; import com.wemirr.framework.commons.exception.CheckedException; +import com.wemirr.platform.ai.core.model.VectorSearchResult; import com.wemirr.platform.ai.domain.dto.resp.EmbeddingMatchResp; import com.wemirr.platform.ai.domain.entity.KnowledgeBase; import com.wemirr.platform.ai.domain.entity.KnowledgeChunk; @@ -10,8 +11,6 @@ import com.wemirr.platform.ai.service.KnowledgeBaseService; import com.wemirr.platform.ai.service.KnowledgeSearchService; import com.wemirr.platform.ai.service.ModelService; import com.wemirr.platform.ai.service.VectorSearchService; -import dev.langchain4j.data.segment.TextSegment; -import dev.langchain4j.store.embedding.EmbeddingMatch; import lombok.RequiredArgsConstructor; import lombok.extern.slf4j.Slf4j; import org.springframework.stereotype.Service; @@ -57,14 +56,14 @@ public class KnowledgeSearchServiceImpl implements KnowledgeSearchService { return null; } - List> matches = vectorSearchService.search(knowledgeBase, embeddingModel, query, topK); + List matches = vectorSearchService.search(knowledgeBase, embeddingModel, query, topK); return matches.stream() .map(match -> { return EmbeddingMatchResp.builder() - .content(match.embedded().text()) - .score(match.score()) - .metadata(match.embedded().metadata().toMap()) + .content(match.getContent()) + .score(match.getScore()) + .metadata(match.getMetadata()) //语义搜索(向量) .searchType("semantic") .build(); @@ -175,14 +174,14 @@ public class KnowledgeSearchServiceImpl implements KnowledgeSearchService { return keywordSearch(kbId, query, topK); } - List> matches = vectorSearchService.search(knowledgeBase, embeddingModel, query, topK); + List matches = vectorSearchService.search(knowledgeBase, embeddingModel, query, topK); return matches.stream() .map(match -> { Map result = new HashMap<>(); - result.put("content", match.embedded().text()); - result.put("score", match.score()); - result.put("metadata", match.embedded().metadata().toMap()); + result.put("content", match.getContent()); + result.put("score", match.getScore()); + result.put("metadata", match.getMetadata()); result.put("searchType", "semantic"); return result; }) diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/service/impl/McpConnectionManagerImpl.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/service/impl/McpConnectionManagerImpl.java index 441e379d..0c2f5343 100644 --- a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/service/impl/McpConnectionManagerImpl.java +++ b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/service/impl/McpConnectionManagerImpl.java @@ -1,29 +1,31 @@ package com.wemirr.platform.ai.service.impl; -import com.fasterxml.jackson.core.type.TypeReference; -import com.fasterxml.jackson.databind.ObjectMapper; import com.wemirr.framework.commons.exception.CheckedException; +import com.wemirr.platform.ai.core.provider.mcp.McpClientHandle; import com.wemirr.platform.ai.domain.entity.McpServer; import com.wemirr.platform.ai.repository.McpServerMapper; import com.wemirr.platform.ai.service.McpConnectionManager; -import dev.langchain4j.mcp.client.DefaultMcpClient; -import dev.langchain4j.mcp.client.McpClient; -import dev.langchain4j.mcp.client.transport.McpTransport; -import dev.langchain4j.mcp.client.transport.http.StreamableHttpMcpTransport; -import dev.langchain4j.mcp.client.transport.stdio.StdioMcpTransport; +import io.modelcontextprotocol.client.McpClient; +import io.modelcontextprotocol.client.McpSyncClient; +import io.modelcontextprotocol.client.transport.HttpClientSseClientTransport; +import io.modelcontextprotocol.client.transport.ServerParameters; +import io.modelcontextprotocol.client.transport.StdioClientTransport; +import io.modelcontextprotocol.json.McpJsonDefaults; +import io.modelcontextprotocol.spec.McpClientTransport; +import io.modelcontextprotocol.spec.McpSchema; import lombok.RequiredArgsConstructor; import lombok.extern.slf4j.Slf4j; import org.springframework.stereotype.Service; +import org.springframework.util.StringUtils; +import java.net.URI; import java.time.Duration; -import java.util.ArrayList; -import java.util.Collections; import java.util.List; import java.util.Map; import java.util.concurrent.ConcurrentHashMap; /** - * MCP连接管理器实现 + * MCP 连接管理器实现。 * * @author xJh * @since 2025/12/07 @@ -33,112 +35,128 @@ import java.util.concurrent.ConcurrentHashMap; @RequiredArgsConstructor public class McpConnectionManagerImpl implements McpConnectionManager { + private static final Duration MCP_TIMEOUT = Duration.ofSeconds(30); + private final McpServerMapper mcpServerMapper; - private final ObjectMapper objectMapper; - - // 缓存客户端实例: configId -> McpClient - private final Map clientCache = new ConcurrentHashMap<>(); + + private final Map clientCache = new ConcurrentHashMap<>(); @Override - public McpClient getClient(Long configId) { - if (clientCache.containsKey(configId)) { - return clientCache.get(configId); + public McpClientHandle getClient(Long configId) { + McpClientHandle cached = clientCache.get(configId); + if (cached != null) { + return cached; } - McpServer config = mcpServerMapper.selectById(configId); if (config == null) { throw CheckedException.notFound("MCP 配置不存在: " + configId); } - if (Boolean.FALSE.equals(config.getStatus())) { throw CheckedException.badRequest("MCP 服务已禁用: " + config.getName()); } - - McpClient client = createClient(config); + McpClientHandle client = createClient(config); clientCache.put(configId, client); return client; } @Override public void closeClient(Long configId) { - McpClient client = clientCache.remove(configId); + McpClientHandle client = clientCache.remove(configId); if (client == null) { return; } try { client.close(); } catch (Exception e) { - log.error("Error closing MCP client: {}", configId, e); + log.error("关闭 MCP 客户端失败: {}", configId, e); } } @Override public void refreshClient(Long configId) { closeClient(configId); - // 下次 getClient 时会自动重新创建 } - private McpClient createClient(McpServer config) { - McpTransport transport; - - if ("STDIO".equalsIgnoreCase(config.getType())) { - Map env = config.getEnv(); - List args = config.getArgs(); - - List fullCommand = new ArrayList<>(); - fullCommand.add(config.getCommand()); - if (args != null && !args.isEmpty()) { - fullCommand.addAll(args); - } - - StdioMcpTransport.Builder builder = StdioMcpTransport.builder() - .command(fullCommand) - .logEvents(true); - - if (env != null && !env.isEmpty()) { - builder.environment(env); + private McpClientHandle createClient(McpServer config) { + McpClientTransport transport = createTransport(config); + McpSyncClient client = McpClient.sync(transport) + .requestTimeout(MCP_TIMEOUT) + .initializationTimeout(MCP_TIMEOUT) + .build(); + try { + client.initialize(); + return new SpringAiMcpClientHandle(config.getId(), client); + } catch (Exception e) { + try { + client.close(); + } catch (Exception closeException) { + log.warn("关闭初始化失败的 MCP 客户端失败: {}", config.getId(), closeException); } - - transport = builder.build(); - } else if ("SSE".equalsIgnoreCase(config.getType()) || "HTTP".equalsIgnoreCase(config.getType())) { - transport = StreamableHttpMcpTransport.builder() - .url(config.getUrl()) - .logRequests(true) - .logResponses(true) - .build(); - } else { - throw CheckedException.badRequest("不支持的 MCP 传输类型: " + config.getType()); + throw CheckedException.badRequest("MCP 服务连接失败: " + config.getName() + ", " + e.getMessage()); + } + } + + private McpClientTransport createTransport(McpServer config) { + String type = config.getType(); + if (!StringUtils.hasText(type)) { + throw CheckedException.badRequest("MCP 连接类型不能为空"); + } + if ("STDIO".equalsIgnoreCase(type)) { + return createStdioTransport(config); + } + if ("SSE".equalsIgnoreCase(type)) { + return createSseTransport(config); } + throw CheckedException.badRequest("不支持的 MCP 连接类型: " + type); + } - return new DefaultMcpClient.Builder() - .key("wemirr-platform-ai-" + config.getId()) - .transport(transport) - .toolExecutionTimeout(Duration.ofSeconds(60)) + private McpClientTransport createStdioTransport(McpServer config) { + if (!StringUtils.hasText(config.getCommand())) { + throw CheckedException.badRequest("STDIO MCP 命令不能为空"); + } + ServerParameters parameters = ServerParameters.builder(config.getCommand()) + .args(config.getArgs() == null ? List.of() : config.getArgs()) + .env(config.getEnv() == null ? Map.of() : config.getEnv()) .build(); + return new StdioClientTransport(parameters, McpJsonDefaults.getMapper()); } - private Map parseEnv(String json) { - if (json == null || json.isEmpty()) { - return Collections.emptyMap(); + private McpClientTransport createSseTransport(McpServer config) { + if (!StringUtils.hasText(config.getUrl())) { + throw CheckedException.badRequest("SSE MCP URL 不能为空"); } - try { - return objectMapper.readValue(json, new TypeReference>() {}); - } catch (Exception e) { - log.error("Failed to parse env JSON", e); - return Collections.emptyMap(); + URI uri = URI.create(config.getUrl()); + String baseUri = uri.getScheme() + "://" + uri.getRawAuthority(); + String endpoint = StringUtils.hasText(uri.getRawPath()) ? uri.getRawPath() : "/sse"; + if (StringUtils.hasText(uri.getRawQuery())) { + endpoint = endpoint + "?" + uri.getRawQuery(); } + return HttpClientSseClientTransport.builder(baseUri) + .sseEndpoint(endpoint) + .connectTimeout(MCP_TIMEOUT) + .build(); } - private List parseArgs(String json) { - if (json == null || json.isEmpty()) { - return Collections.emptyList(); + private record SpringAiMcpClientHandle(Long configId, McpSyncClient client) implements McpClientHandle { + + @Override + public List listTools() { + try { + McpSchema.ListToolsResult result = client.listTools(); + if (result == null || result.tools() == null) { + return List.of(); + } + return result.tools().stream() + .map(tool -> new McpToolDescriptor(tool.name(), tool.description())) + .toList(); + } catch (Exception e) { + throw CheckedException.badRequest("获取 MCP 工具列表失败: " + configId + ", " + e.getMessage()); + } } - try { - return objectMapper.readValue(json, new TypeReference>() {}); - } catch (Exception e) { - log.error("Failed to parse args JSON", e); - return Collections.emptyList(); + + @Override + public void close() { + client.close(); } } } - diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/service/impl/McpServerConfigServiceImpl.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/service/impl/McpServerConfigServiceImpl.java index 84239dcc..df931ecd 100644 --- a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/service/impl/McpServerConfigServiceImpl.java +++ b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/service/impl/McpServerConfigServiceImpl.java @@ -12,10 +12,9 @@ import com.wemirr.platform.ai.domain.dto.resp.McpConnectionTestResp; import com.wemirr.platform.ai.domain.dto.resp.McpToolInfoResp; import com.wemirr.platform.ai.domain.entity.McpServer; import com.wemirr.platform.ai.repository.McpServerMapper; +import com.wemirr.platform.ai.core.provider.mcp.McpClientHandle; import com.wemirr.platform.ai.service.McpConnectionManager; import com.wemirr.platform.ai.service.McpServerConfigService; -import dev.langchain4j.agent.tool.ToolSpecification; -import dev.langchain4j.mcp.client.McpClient; import lombok.RequiredArgsConstructor; import lombok.extern.slf4j.Slf4j; import org.springframework.stereotype.Service; @@ -88,10 +87,10 @@ public class McpServerConfigServiceImpl extends SuperServiceImpl tools = client.listTools(); + List tools = client.listTools(); long responseTime = System.currentTimeMillis() - startTime; @@ -119,15 +118,15 @@ public class McpServerConfigServiceImpl extends SuperServiceImpl getTools(Long id) { - McpClient client = mcpConnectionManager.getClient(id); - List toolSpecs = client.listTools(); + McpClientHandle client = mcpConnectionManager.getClient(id); + List toolSpecs = client.listTools(); if (toolSpecs == null || toolSpecs.isEmpty()) { return new ArrayList<>(); } List result = new ArrayList<>(); - for (ToolSpecification spec : toolSpecs) { + for (McpClientHandle.McpToolDescriptor spec : toolSpecs) { McpToolInfoResp toolInfo = McpToolInfoResp.builder() .name(spec.name()) .description(spec.description()) diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/service/impl/VectorSearchServiceImpl.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/service/impl/VectorSearchServiceImpl.java index aab80201..cb0a6e0e 100644 --- a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/service/impl/VectorSearchServiceImpl.java +++ b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/service/impl/VectorSearchServiceImpl.java @@ -1,18 +1,15 @@ package com.wemirr.platform.ai.service.impl; -import com.wemirr.platform.ai.core.provider.embedding.EmbeddingModelService; +import com.wemirr.platform.ai.core.model.VectorSearchResult; import com.wemirr.platform.ai.core.provider.vector.VectorStoreFactory; import com.wemirr.platform.ai.domain.entity.KnowledgeBase; import com.wemirr.platform.ai.domain.entity.ModelEntity; import com.wemirr.platform.ai.service.VectorSearchService; -import dev.langchain4j.data.embedding.Embedding; -import dev.langchain4j.data.segment.TextSegment; -import dev.langchain4j.model.embedding.EmbeddingModel; -import dev.langchain4j.store.embedding.EmbeddingMatch; -import dev.langchain4j.store.embedding.EmbeddingSearchRequest; -import dev.langchain4j.store.embedding.EmbeddingStore; import lombok.RequiredArgsConstructor; import lombok.extern.slf4j.Slf4j; +import org.springframework.ai.document.Document; +import org.springframework.ai.vectorstore.SearchRequest; +import org.springframework.ai.vectorstore.VectorStore; import org.springframework.stereotype.Service; import java.util.List; @@ -30,33 +27,29 @@ import java.util.List; public class VectorSearchServiceImpl implements VectorSearchService { private final VectorStoreFactory vectorStoreFactory; - private final EmbeddingModelService embeddingModelService; @Override - public List> search(KnowledgeBase knowledgeBase, ModelEntity modelEntity, String query, int topK) { + public List search(KnowledgeBase knowledgeBase, ModelEntity modelEntity, String query, int topK) { try { - // 1. 获取向量存储 - EmbeddingStore embeddingStore = vectorStoreFactory.createForKnowledgeBase(knowledgeBase, modelEntity); - - // 2. 创建嵌入模型实例 - EmbeddingModel embeddingModel = embeddingModelService.getModel(modelEntity); - - // 3. 生成查询向量 - Embedding queryEmbedding = embeddingModel.embed(query).content(); - - // 4. 创建搜索请求 - EmbeddingSearchRequest embeddingSearchRequest = EmbeddingSearchRequest.builder() - .queryEmbedding(queryEmbedding) - .maxResults(topK) + VectorStore vectorStore = vectorStoreFactory.createForKnowledgeBase(knowledgeBase, modelEntity); + SearchRequest searchRequest = SearchRequest.builder() + .query(query) + .topK(topK) + .similarityThresholdAll() .build(); - - // 5. 执行向量搜索 - List> matches = embeddingStore.search(embeddingSearchRequest).matches(); + List documents = vectorStore.similaritySearch(searchRequest); log.info("向量搜索完成: kbId={}, query={}, topK={}, results={}", - knowledgeBase.getId(), query, topK, matches.size()); + knowledgeBase.getId(), query, topK, documents.size()); - return matches; + return documents.stream() + .map(document -> VectorSearchResult.builder() + .id(document.getId()) + .content(document.getText()) + .score(document.getScore()) + .metadata(document.getMetadata()) + .build()) + .toList(); } catch (Exception e) { log.error("向量搜索失败: kbId={}, query={}", knowledgeBase.getId(), query, e); @@ -70,9 +63,6 @@ public class VectorSearchServiceImpl implements VectorSearchService { // 尝试创建向量存储连接 vectorStoreFactory.createForKnowledgeBase(knowledgeBase, modelEntity); - // 尝试创建嵌入模型 - embeddingModelService.getModel(modelEntity); - log.debug("向量存储可用性检查通过: kbId={}, model={}", knowledgeBase.getId(), modelEntity.getName()); return true; diff --git a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/service/impl/WorkflowExecutionServiceImpl.java b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/service/impl/WorkflowExecutionServiceImpl.java index 80f8885c..0d83f8f5 100644 --- a/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/service/impl/WorkflowExecutionServiceImpl.java +++ b/wemirr-plugin/wemirr-platform-ai/src/main/java/com/wemirr/platform/ai/service/impl/WorkflowExecutionServiceImpl.java @@ -121,7 +121,7 @@ public class WorkflowExecutionServiceImpl extends SuperServiceImpl nodes, List edges) { diff --git a/wemirr-plugin/wemirr-platform-ai/src/test/java/com/wemirr/platform/ai/integration/ParameterExtractorNodeExecutorIntegrationTest.java b/wemirr-plugin/wemirr-platform-ai/src/test/java/com/wemirr/platform/ai/integration/ParameterExtractorNodeExecutorIntegrationTest.java index cc36dfd4..ffb11273 100644 --- a/wemirr-plugin/wemirr-platform-ai/src/test/java/com/wemirr/platform/ai/integration/ParameterExtractorNodeExecutorIntegrationTest.java +++ b/wemirr-plugin/wemirr-platform-ai/src/test/java/com/wemirr/platform/ai/integration/ParameterExtractorNodeExecutorIntegrationTest.java @@ -27,10 +27,11 @@ import com.wemirr.platform.ai.core.workflow.agent.impl.ParameterExtractorNodeExe import com.wemirr.platform.ai.core.workflow.context.ExecutionContext; import com.wemirr.platform.ai.domain.entity.ModelEntity; import com.wemirr.platform.ai.domain.entity.workflow.WorkflowNode; -import dev.langchain4j.data.message.AiMessage; -import dev.langchain4j.model.chat.ChatModel; -import dev.langchain4j.model.chat.response.ChatResponse; -import dev.langchain4j.model.output.TokenUsage; +import org.springframework.ai.chat.messages.AssistantMessage; +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.prompt.Prompt; import org.junit.jupiter.api.*; import org.junit.jupiter.api.extension.ExtendWith; import org.mockito.Mock; @@ -106,10 +107,10 @@ class ParameterExtractorNodeExecutorIntegrationTest { lenient().when(textModelService.model(any(ModelEntity.class))).thenReturn(chatModel); ChatResponse chatResponse = mock(ChatResponse.class); - AiMessage aiMessage = AiMessage.from(responseContent); - lenient().when(chatResponse.aiMessage()).thenReturn(aiMessage); - lenient().when(chatResponse.tokenUsage()).thenReturn(new TokenUsage(20, 10, 30)); - lenient().when(chatModel.chat(anyList())).thenReturn(chatResponse); + Generation generation = mock(Generation.class); + lenient().when(generation.getOutput()).thenReturn(new AssistantMessage(responseContent)); + lenient().when(chatResponse.getResult()).thenReturn(generation); + lenient().when(chatModel.call(any(Prompt.class))).thenReturn(chatResponse); } private List> createBookingParameters() { @@ -334,7 +335,7 @@ class ParameterExtractorNodeExecutorIntegrationTest { agent.execute(node, context); // Assert - verify(chatModel).chat(anyList()); + verify(chatModel).call(any(Prompt.class)); } } @@ -546,7 +547,7 @@ class ParameterExtractorNodeExecutorIntegrationTest { // Assert - Second call assertTrue(result2.isSuccess()); // 验证 chatModel 被调用了两次 - verify(chatModel, times(2)).chat(anyList()); + verify(chatModel, times(2)).call(any(Prompt.class)); } } diff --git a/wemirr-plugin/wemirr-platform-ai/src/test/java/com/wemirr/platform/ai/integration/QuestionClassifierNodeExecutorIntegrationTest.java b/wemirr-plugin/wemirr-platform-ai/src/test/java/com/wemirr/platform/ai/integration/QuestionClassifierNodeExecutorIntegrationTest.java index d4b0b5de..b1cc205d 100644 --- a/wemirr-plugin/wemirr-platform-ai/src/test/java/com/wemirr/platform/ai/integration/QuestionClassifierNodeExecutorIntegrationTest.java +++ b/wemirr-plugin/wemirr-platform-ai/src/test/java/com/wemirr/platform/ai/integration/QuestionClassifierNodeExecutorIntegrationTest.java @@ -29,10 +29,13 @@ import com.wemirr.platform.ai.core.workflow.config.node.QuestionClassifierConfig import com.wemirr.platform.ai.core.workflow.context.ExecutionContext; import com.wemirr.platform.ai.domain.entity.ModelEntity; import com.wemirr.platform.ai.domain.entity.workflow.WorkflowNode; -import dev.langchain4j.data.message.AiMessage; -import dev.langchain4j.model.chat.ChatModel; -import dev.langchain4j.model.chat.response.ChatResponse; -import dev.langchain4j.model.output.TokenUsage; +import org.springframework.ai.chat.messages.AssistantMessage; +import org.springframework.ai.chat.metadata.ChatResponseMetadata; +import org.springframework.ai.chat.metadata.Usage; +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.prompt.Prompt; import org.junit.jupiter.api.*; import org.junit.jupiter.api.extension.ExtendWith; import org.mockito.Mock; @@ -102,10 +105,17 @@ class QuestionClassifierNodeExecutorIntegrationTest { lenient().when(textModelService.model(any(ModelEntity.class))).thenReturn(chatModel); ChatResponse chatResponse = mock(ChatResponse.class); - AiMessage aiMessage = AiMessage.from(responseContent); - lenient().when(chatResponse.aiMessage()).thenReturn(aiMessage); - lenient().when(chatResponse.tokenUsage()).thenReturn(new TokenUsage(10, 5, 15)); - lenient().when(chatModel.chat(anyList())).thenReturn(chatResponse); + Generation generation = mock(Generation.class); + ChatResponseMetadata metadata = mock(ChatResponseMetadata.class); + Usage usage = mock(Usage.class); + lenient().when(generation.getOutput()).thenReturn(new AssistantMessage(responseContent)); + lenient().when(chatResponse.getResult()).thenReturn(generation); + lenient().when(chatResponse.getMetadata()).thenReturn(metadata); + lenient().when(metadata.getUsage()).thenReturn(usage); + lenient().when(usage.getPromptTokens()).thenReturn(10); + lenient().when(usage.getCompletionTokens()).thenReturn(5); + lenient().when(usage.getTotalTokens()).thenReturn(15); + lenient().when(chatModel.call(any(Prompt.class))).thenReturn(chatResponse); } private List> createTechSupportCategories() { @@ -161,7 +171,7 @@ class QuestionClassifierNodeExecutorIntegrationTest { assertEquals("technical", result.getOutputs().get("category")); assertEquals("Technical Support", result.getOutputs().get("categoryName")); assertEquals("technical", result.getNextBranch()); - verify(chatModel).chat(anyList()); + verify(chatModel).call(any(Prompt.class)); } @Test diff --git "a/\351\231\204\344\273\266/mysql/v4-ai.sql" "b/\351\231\204\344\273\266/mysql/v4-ai.sql" index 95b6ef06..27dec287 100644 --- "a/\351\231\204\344\273\266/mysql/v4-ai.sql" +++ "b/\351\231\204\344\273\266/mysql/v4-ai.sql" @@ -33,7 +33,7 @@ CREATE TABLE `ai_agent` ( `tools` varchar(500) CHARACTER SET utf8mb4 COLLATE utf8mb4_0900_ai_ci DEFAULT NULL COMMENT '工具/函数配置', `tenant_id` bigint DEFAULT NULL COMMENT '租户ID', `avatar` varchar(255) CHARACTER SET utf8mb4 COLLATE utf8mb4_0900_ai_ci DEFAULT NULL COMMENT 'Agent 头像', - `system_prompt` varchar(255) DEFAULT NULL COMMENT '预设系统提示词', + `system_prompt` varchar(2550) DEFAULT NULL COMMENT '预设系统提示词', `mcp_server_ids` varchar(255) DEFAULT NULL, `deleted` tinyint(1) DEFAULT '0' COMMENT '逻辑删除', `create_by` bigint DEFAULT NULL COMMENT '创建人ID', -- Gitee