This is an automated email from the ASF dual-hosted git repository. tomsun28 pushed a commit to branch 2.0.0 in repository https://gitbox.apache.org/repos/asf/hertzbeat.git
commit b58913824b5876c16422e80f5fcb374c5b7ab491 Author: tomsun28 <[email protected]> AuthorDate: Fri Oct 9 16:17:24 2026 +0800 fix(ai): restore provider compatibility and secret decryption ordering DeepSeek-compatible providers reject tool names containing dots, so encode names on the wire (monitor.get -> monitor_get) and decode provider responses back to the canonical registry names, covering tool definitions, history tool calls, and tool responses. Also fix a startup ordering bug where ReloadableAgentRuntimeModelClient reloaded in its constructor, before ConfigInitializer installed the AES secret key, so persisted apiKey ciphertext failed to decrypt and was sent to the provider verbatim (401). Re-run the reload in SmartLifecycle start() after secret initialization. Harden the secret pipeline so an undecryptable payload fails loudly instead of leaking ciphertext, and so re-saving never double-encrypts a corrupted key. Generated with [Devin](https://devin.ai) Co-Authored-By: Devin <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../ai/gateway/runtime/HertzBeatModel.java | 110 ++++++++++++++++----- .../runtime/ReloadableAgentRuntimeModelClient.java | 41 +++++++- .../ai/gateway/runtime/HertzBeatModelTest.java | 25 ++++- .../ReloadableAgentRuntimeModelClientTest.java | 29 ++++++ .../org/apache/hertzbeat/common/util/AesUtil.java | 21 ++++ .../impl/ModelProviderConfigServiceImpl.java | 7 +- .../impl/ModelProviderConfigServiceImplTest.java | 16 +++ 7 files changed, 219 insertions(+), 30 deletions(-) diff --git a/hertzbeat-ai-gateway/src/main/java/org/apache/hertzbeat/ai/gateway/runtime/HertzBeatModel.java b/hertzbeat-ai-gateway/src/main/java/org/apache/hertzbeat/ai/gateway/runtime/HertzBeatModel.java index 1ff628b4ab..fd60891bbf 100644 --- a/hertzbeat-ai-gateway/src/main/java/org/apache/hertzbeat/ai/gateway/runtime/HertzBeatModel.java +++ b/hertzbeat-ai-gateway/src/main/java/org/apache/hertzbeat/ai/gateway/runtime/HertzBeatModel.java @@ -28,6 +28,7 @@ import java.util.Map; import java.util.Objects; import java.util.UUID; import java.util.function.Consumer; +import java.util.regex.Pattern; import org.apache.hertzbeat.ai.gateway.runtime.provider.AgentModelRequestOptionsFactory; import org.apache.hertzbeat.ai.gateway.tool.core.AgentToolDescriptor; import org.springframework.ai.chat.messages.AssistantMessage; @@ -54,6 +55,7 @@ public class HertzBeatModel { private static final int MODEL_ERROR_MESSAGE_LIMIT = 1024; private static final String FUNCTION_TOOL_TYPE = "function"; + private static final Pattern WIRE_UNSAFE_TOOL_NAME_CHARS = Pattern.compile("[^a-zA-Z0-9_-]"); private static final TypeReference<Map<String, Object>> MAP_TYPE = new TypeReference<>() { }; @@ -84,10 +86,11 @@ public class HertzBeatModel { public AgentRuntimeModelResponse stream(AgentRuntimeModelRequest request, AgentRuntimeControl control, Consumer<String> textDeltaConsumer) { control.checkpoint(); - Prompt prompt = toPrompt(request); + ToolNameCodec toolNameCodec = ToolNameCodec.of(request.getAvailableTools()); + Prompt prompt = toPrompt(request, toolNameCodec); Thread currentThread = Thread.currentThread(); AutoCloseable abortRegistration = control.onAbort(currentThread::interrupt); - ChatResponseAccumulator accumulator = new ChatResponseAccumulator(textDeltaConsumer); + ChatResponseAccumulator accumulator = new ChatResponseAccumulator(textDeltaConsumer, toolNameCodec); try { chatModel.stream(prompt) .doOnNext(response -> { @@ -117,14 +120,14 @@ public class HertzBeatModel { return accumulator.toRuntimeResponse(); } - private Prompt toPrompt(AgentRuntimeModelRequest request) { + private Prompt toPrompt(AgentRuntimeModelRequest request, ToolNameCodec toolNameCodec) { RuntimePrompt runtimePrompt = request.getPrompt(); List<Message> messages = new ArrayList<>(); - List<ToolCallback> toolCallbacks = toToolCallbacks(request.getAvailableTools()); + List<ToolCallback> toolCallbacks = toToolCallbacks(request.getAvailableTools(), toolNameCodec); addBaseInstructions(messages, runtimePrompt); addPromptBlocks(messages, runtimePrompt, RuntimePrompt.Role.SYSTEM); addPromptBlocks(messages, runtimePrompt, RuntimePrompt.Role.USER); - addHistoryMessages(messages, request.getChatHistory()); + addHistoryMessages(messages, request.getChatHistory(), toolNameCodec); ChatOptions options = requestOptionsFactory.create(request, toolCallbacks); return new Prompt(messages, options); } @@ -166,19 +169,20 @@ public class HertzBeatModel { .build(); } - private void addHistoryMessages(List<Message> messages, List<TranscriptMessage> chatHistory) { + private void addHistoryMessages(List<Message> messages, List<TranscriptMessage> chatHistory, + ToolNameCodec toolNameCodec) { if (chatHistory.isEmpty()) { return; } for (TranscriptMessage historyMessage : chatHistory) { - Message message = toSpringHistoryMessage(historyMessage); + Message message = toSpringHistoryMessage(historyMessage, toolNameCodec); if (message != null) { messages.add(message); } } } - private Message toSpringHistoryMessage(TranscriptMessage historyMessage) { + private Message toSpringHistoryMessage(TranscriptMessage historyMessage, ToolNameCodec toolNameCodec) { TranscriptMessage.TranscriptRole role = historyMessage.getRole(); if (role == TranscriptMessage.TranscriptRole.USER) { return UserMessage.builder() @@ -193,20 +197,21 @@ public class HertzBeatModel { .build(); } if (role == TranscriptMessage.TranscriptRole.ASSISTANT && !historyMessage.toolCalls().isEmpty()) { - return assistantToolCallHistoryMessage(historyMessage); + return assistantToolCallHistoryMessage(historyMessage, toolNameCodec); } if (role == TranscriptMessage.TranscriptRole.ASSISTANT) { return assistantTextHistoryMessage(historyMessage); } if (role == TranscriptMessage.TranscriptRole.TOOL_RESULT) { - return toolResponseHistoryMessage(historyMessage); + return toolResponseHistoryMessage(historyMessage, toolNameCodec); } return null; } - private AssistantMessage assistantToolCallHistoryMessage(TranscriptMessage historyMessage) { + private AssistantMessage assistantToolCallHistoryMessage(TranscriptMessage historyMessage, + ToolNameCodec toolNameCodec) { List<AssistantMessage.ToolCall> toolCalls = historyMessage.toolCalls().stream() - .map(this::springToolCall) + .map(block -> springToolCall(block, toolNameCodec)) .toList(); return AssistantMessage.builder() .content(historyMessage.text()) @@ -215,11 +220,11 @@ public class HertzBeatModel { .build(); } - private AssistantMessage.ToolCall springToolCall(TranscriptContent block) { + private AssistantMessage.ToolCall springToolCall(TranscriptContent block, ToolNameCodec toolNameCodec) { return new AssistantMessage.ToolCall( block.getId(), FUNCTION_TOOL_TYPE, - block.getName(), + toolNameCodec.encode(block.getName()), assistantToolArguments(block)); } @@ -230,10 +235,11 @@ public class HertzBeatModel { .build(); } - private ToolResponseMessage toolResponseHistoryMessage(TranscriptMessage historyMessage) { + private ToolResponseMessage toolResponseHistoryMessage(TranscriptMessage historyMessage, + ToolNameCodec toolNameCodec) { ToolResponseMessage.ToolResponse response = new ToolResponseMessage.ToolResponse( historyMessage.getToolCallId(), - historyMessage.getToolName(), + toolNameCodec.encode(historyMessage.getToolName()), toolResponseData(historyMessage)); return ToolResponseMessage.builder() .responses(List.of(response)) @@ -329,7 +335,8 @@ public class HertzBeatModel { "Runtime model returned neither a final answer nor tool calls.", usage); } - private List<AgentRuntimeToolCall> toRuntimeToolCalls(List<AssistantMessage.ToolCall> toolCalls) { + private List<AgentRuntimeToolCall> toRuntimeToolCalls(List<AssistantMessage.ToolCall> toolCalls, + ToolNameCodec toolNameCodec) { // Spring AI returns null when the assistant message has no tool calls. if (toolCalls == null || toolCalls.isEmpty()) { return List.of(); @@ -344,20 +351,21 @@ public class HertzBeatModel { } result.add(AgentRuntimeToolCall.builder() .toolCallId(toolCallId) - .toolName(toolCall.name()) + .toolName(toolNameCodec.decode(toolCall.name())) .arguments(arguments) .build()); } return List.copyOf(result); } - private List<ToolCallback> toToolCallbacks(List<AgentToolDescriptor> availableTools) { + private List<ToolCallback> toToolCallbacks(List<AgentToolDescriptor> availableTools, + ToolNameCodec toolNameCodec) { if (availableTools.isEmpty()) { return List.of(); } List<ToolCallback> callbacks = new ArrayList<>(availableTools.size()); for (AgentToolDescriptor tool : availableTools) { - callbacks.add(new DisabledRuntimeToolCallback(tool)); + callbacks.add(new DisabledRuntimeToolCallback(tool, toolNameCodec)); } return List.copyOf(callbacks); } @@ -407,12 +415,14 @@ public class HertzBeatModel { private final class ChatResponseAccumulator { private final Consumer<String> textDeltaConsumer; + private final ToolNameCodec toolNameCodec; private final StringBuilder text = new StringBuilder(); private List<AgentRuntimeToolCall> toolCalls = List.of(); private ChatResponseMetadata metadata; - private ChatResponseAccumulator(Consumer<String> textDeltaConsumer) { + private ChatResponseAccumulator(Consumer<String> textDeltaConsumer, ToolNameCodec toolNameCodec) { this.textDeltaConsumer = textDeltaConsumer; + this.toolNameCodec = toolNameCodec; } private void accept(ChatResponse response) { @@ -427,7 +437,7 @@ public class HertzBeatModel { return; } AssistantMessage output = generation.getOutput(); - List<AgentRuntimeToolCall> responseToolCalls = toRuntimeToolCalls(output.getToolCalls()); + List<AgentRuntimeToolCall> responseToolCalls = toRuntimeToolCalls(output.getToolCalls(), toolNameCodec); if (!responseToolCalls.isEmpty()) { toolCalls = responseToolCalls; } @@ -451,9 +461,9 @@ public class HertzBeatModel { private final ToolDefinition toolDefinition; - private DisabledRuntimeToolCallback(AgentToolDescriptor tool) { + private DisabledRuntimeToolCallback(AgentToolDescriptor tool, ToolNameCodec toolNameCodec) { this.toolDefinition = ToolDefinition.builder() - .name(tool.getName()) + .name(toolNameCodec.encode(tool.getName())) .description(tool.getDescription()) .inputSchema(tool.getInputSchema()) .build(); @@ -469,4 +479,56 @@ public class HertzBeatModel { throw new UnsupportedOperationException("Tool execution is owned by AgentRuntimeLoop"); } } + + /** + * Maps canonical tool names (namespaced with '.') to wire-safe names accepted by providers that enforce + * the OpenAI function-name pattern {@code ^[a-zA-Z0-9_-]+$}, and maps provider tool call names back. + */ + private static final class ToolNameCodec { + + private final Map<String, String> canonicalToWire; + private final Map<String, String> wireToCanonical; + + private ToolNameCodec(Map<String, String> canonicalToWire, Map<String, String> wireToCanonical) { + this.canonicalToWire = canonicalToWire; + this.wireToCanonical = wireToCanonical; + } + + private static ToolNameCodec of(List<AgentToolDescriptor> tools) { + Map<String, String> canonicalToWire = new LinkedHashMap<>(); + Map<String, String> wireToCanonical = new LinkedHashMap<>(); + if (tools == null) { + return new ToolNameCodec(canonicalToWire, wireToCanonical); + } + for (AgentToolDescriptor tool : tools) { + String canonical = tool.getName(); + if (!StringUtils.hasText(canonical)) { + continue; + } + String wire = sanitize(canonical); + int suffix = 2; + while (wireToCanonical.containsKey(wire) && !canonical.equals(wireToCanonical.get(wire))) { + wire = sanitize(canonical) + "_" + suffix++; + } + canonicalToWire.put(canonical, wire); + wireToCanonical.put(wire, canonical); + } + return new ToolNameCodec(canonicalToWire, wireToCanonical); + } + + private static String sanitize(String name) { + return name == null ? null : WIRE_UNSAFE_TOOL_NAME_CHARS.matcher(name).replaceAll("_"); + } + + private String encode(String name) { + if (name == null) { + return null; + } + return canonicalToWire.getOrDefault(name, sanitize(name)); + } + + private String decode(String name) { + return wireToCanonical.getOrDefault(name, name); + } + } } diff --git a/hertzbeat-ai-gateway/src/main/java/org/apache/hertzbeat/ai/gateway/runtime/ReloadableAgentRuntimeModelClient.java b/hertzbeat-ai-gateway/src/main/java/org/apache/hertzbeat/ai/gateway/runtime/ReloadableAgentRuntimeModelClient.java index 8626de5277..46126d5107 100644 --- a/hertzbeat-ai-gateway/src/main/java/org/apache/hertzbeat/ai/gateway/runtime/ReloadableAgentRuntimeModelClient.java +++ b/hertzbeat-ai-gateway/src/main/java/org/apache/hertzbeat/ai/gateway/runtime/ReloadableAgentRuntimeModelClient.java @@ -17,6 +17,7 @@ package org.apache.hertzbeat.ai.gateway.runtime; +import java.util.concurrent.atomic.AtomicBoolean; import java.util.concurrent.atomic.AtomicReference; import java.util.function.Consumer; import lombok.extern.slf4j.Slf4j; @@ -24,6 +25,8 @@ import org.apache.hertzbeat.ai.gateway.runtime.provider.AgentModelProviderRegist import org.apache.hertzbeat.alert.service.AgentClientAvailability; import org.apache.hertzbeat.common.entity.dto.ModelProviderConfig; import org.apache.hertzbeat.manager.service.ModelProviderConfigurationService; +import org.springframework.context.SmartLifecycle; +import org.springframework.core.Ordered; import org.springframework.stereotype.Component; import org.springframework.util.StringUtils; @@ -32,14 +35,19 @@ import org.springframework.util.StringUtils; */ @Slf4j @Component -public class ReloadableAgentRuntimeModelClient implements AgentRuntimeModelClient, AgentClientAvailability { +public class ReloadableAgentRuntimeModelClient implements AgentRuntimeModelClient, AgentClientAvailability, SmartLifecycle { private static final String DEFAULT_PROVIDER_TYPE = "openai-compatible"; + // Must run after ConfigInitializer (SmartLifecycle, phase Ordered.HIGHEST_PRECEDENCE) so that + // AesUtil's default secret key is installed before persisted provider secrets are decrypted. + private static final int LIFECYCLE_PHASE = Ordered.HIGHEST_PRECEDENCE + 10; + private final ModelProviderConfigurationService configurationService; private final AgentProviderProperties providerProperties; private final AgentModelProviderRegistry providerRegistry; private final AtomicReference<HertzBeatModel> model = new AtomicReference<>(); + private final AtomicBoolean running = new AtomicBoolean(); public ReloadableAgentRuntimeModelClient(ModelProviderConfigurationService configurationService, AgentProviderProperties providerProperties, @@ -73,6 +81,37 @@ public class ReloadableAgentRuntimeModelClient implements AgentRuntimeModelClien * Refresh the runtime after the configuration service has committed a state change. */ public void refreshConfiguration() { + reloadQuietly(); + } + + @Override + public void start() { + // Re-run after all SmartLifecycle initializers so encryption secrets are loaded. + reloadQuietly(); + running.set(true); + } + + @Override + public void stop() { + running.set(false); + } + + @Override + public boolean isRunning() { + return running.get(); + } + + @Override + public boolean isAutoStartup() { + return true; + } + + @Override + public int getPhase() { + return LIFECYCLE_PHASE; + } + + private void reloadQuietly() { try { reload(); } catch (RuntimeException exception) { diff --git a/hertzbeat-ai-gateway/src/test/java/org/apache/hertzbeat/ai/gateway/runtime/HertzBeatModelTest.java b/hertzbeat-ai-gateway/src/test/java/org/apache/hertzbeat/ai/gateway/runtime/HertzBeatModelTest.java index 58cc3c2692..02cdff86c5 100644 --- a/hertzbeat-ai-gateway/src/test/java/org/apache/hertzbeat/ai/gateway/runtime/HertzBeatModelTest.java +++ b/hertzbeat-ai-gateway/src/test/java/org/apache/hertzbeat/ai/gateway/runtime/HertzBeatModelTest.java @@ -183,6 +183,25 @@ class HertzBeatModelTest { assertFalse(chatModel.prompt.getOptions() instanceof ToolCallingChatOptions); } + @Test + void wireSafeToolCallNameShouldDecodeToCanonicalName() { + AssistantMessage.ToolCall toolCall = new AssistantMessage.ToolCall( + "call-9", "function", "monitor_get", "{\"pageSize\":1}"); + AssistantMessage assistantMessage = AssistantMessage.builder() + .content("") + .toolCalls(List.of(toolCall)) + .build(); + CapturingChatModel chatModel = new CapturingChatModel(response(assistantMessage, + ChatResponseMetadata.builder().build())); + HertzBeatModel client = new HertzBeatModel(chatModel); + + AgentRuntimeModelResponse response = stream(client, requestWithTools()); + + assertEquals(AgentRuntimeModelResponse.ResponseType.TOOL_CALLS, response.getType()); + assertEquals(1, response.getToolCalls().size()); + assertEquals("monitor.get", response.getToolCalls().get(0).getToolName()); + } + @Test void missingToolCallIdsShouldBeGenerated() { AssistantMessage assistantMessage = AssistantMessage.builder() @@ -239,7 +258,7 @@ class HertzBeatModelTest { assertEquals(1, options.getToolCallbacks().size()); ToolCallback callback = options.getToolCallbacks().get(0); ToolDefinition definition = callback.getToolDefinition(); - assertEquals("monitor.get", definition.name()); + assertEquals("monitor_get", definition.name()); assertTrue(definition.description().contains("Query monitor inventory")); assertTrue(definition.description().contains("apiKey=secret")); assertTrue(definition.inputSchema().contains("\"type\": \"object\"")); @@ -291,7 +310,7 @@ class HertzBeatModelTest { assertEquals(1, toolCallMessage.getToolCalls().size()); AssistantMessage.ToolCall toolCall = toolCallMessage.getToolCalls().get(0); assertEquals("call-1", toolCall.id()); - assertEquals("alert.history", toolCall.name()); + assertEquals("alert_history", toolCall.name()); assertTrue(toolCall.arguments().contains("\"alertId\":1001") || toolCall.arguments().contains("\"alertId\": 1001")); assertFalse(toolCall.arguments().contains("agc-1")); @@ -300,7 +319,7 @@ class HertzBeatModelTest { assertEquals(1, toolResponseMessage.getResponses().size()); ToolResponseMessage.ToolResponse toolResponse = toolResponseMessage.getResponses().get(0); assertEquals("call-1", toolResponse.id()); - assertEquals("alert.history", toolResponse.name()); + assertEquals("alert_history", toolResponse.name()); assertTrue(toolResponse.responseData().contains("status=SUCCEEDED")); assertFalse(toolResponse.responseData().contains("object://agent-output/1")); diff --git a/hertzbeat-ai-gateway/src/test/java/org/apache/hertzbeat/ai/gateway/runtime/ReloadableAgentRuntimeModelClientTest.java b/hertzbeat-ai-gateway/src/test/java/org/apache/hertzbeat/ai/gateway/runtime/ReloadableAgentRuntimeModelClientTest.java index ebfc83ef78..f62886ef53 100644 --- a/hertzbeat-ai-gateway/src/test/java/org/apache/hertzbeat/ai/gateway/runtime/ReloadableAgentRuntimeModelClientTest.java +++ b/hertzbeat-ai-gateway/src/test/java/org/apache/hertzbeat/ai/gateway/runtime/ReloadableAgentRuntimeModelClientTest.java @@ -22,6 +22,7 @@ import static org.junit.jupiter.api.Assertions.assertFalse; import static org.junit.jupiter.api.Assertions.assertNotNull; import static org.junit.jupiter.api.Assertions.assertNotSame; import static org.junit.jupiter.api.Assertions.assertTrue; +import static org.mockito.Mockito.doReturn; import static org.mockito.Mockito.mock; import static org.mockito.Mockito.when; @@ -34,6 +35,7 @@ import org.apache.hertzbeat.common.entity.dto.ModelProviderConfig; import org.apache.hertzbeat.manager.service.ModelProviderConfigurationService; import org.junit.jupiter.api.Test; import org.springframework.ai.chat.model.ChatModel; +import org.springframework.core.Ordered; /** * Test case for {@link ReloadableAgentRuntimeModelClient}. @@ -91,6 +93,33 @@ class ReloadableAgentRuntimeModelClientTest { assertTrue(client.isAgentClientConfigured()); } + @Test + void startShouldRecoverWhenSecretsWereNotReadyAtConstruction() { + ModelProviderConfigurationService configurationService = mock(ModelProviderConfigurationService.class); + when(configurationService.getActiveConfiguration()) + .thenThrow(new IllegalStateException("secrets not initialized yet")); + TestAgentModelProvider provider = new TestAgentModelProvider(); + ReloadableAgentRuntimeModelClient client = + new ReloadableAgentRuntimeModelClient(configurationService, new AgentProviderProperties(), + new AgentModelProviderRegistry(List.of(provider))); + + assertFalse(client.isAgentClientConfigured()); + + ModelProviderConfig databaseProvider = new ModelProviderConfig(); + databaseProvider.setType("test-provider"); + databaseProvider.setCode("database-preset"); + databaseProvider.setModel("database-model"); + databaseProvider.setApiKey("database-secret"); + doReturn(databaseProvider).when(configurationService).getActiveConfiguration(); + + client.start(); + + assertTrue(client.isRunning()); + assertTrue(client.isAgentClientConfigured()); + assertEquals("database-model", provider.lastCreatedModel().model); + assertTrue(client.getPhase() > Ordered.HIGHEST_PRECEDENCE); + } + @Test void removingDatabaseProviderShouldRestorePropertyProvider() { ModelProviderConfigurationService configurationService = mock(ModelProviderConfigurationService.class); diff --git a/hertzbeat-common-core/src/main/java/org/apache/hertzbeat/common/util/AesUtil.java b/hertzbeat-common-core/src/main/java/org/apache/hertzbeat/common/util/AesUtil.java index 81a65af8ad..1ff09e9258 100644 --- a/hertzbeat-common-core/src/main/java/org/apache/hertzbeat/common/util/AesUtil.java +++ b/hertzbeat-common-core/src/main/java/org/apache/hertzbeat/common/util/AesUtil.java @@ -208,6 +208,27 @@ public final class AesUtil { return true; } + /** + * Determine whether the value carries the HertzBeat encrypted payload header + * without attempting to decrypt it. A payload that matches this shape but + * cannot be decrypted (for example a mismatched secret key) is still + * ciphertext and must never be treated as plaintext. + * @param text text + * @return true false + */ + public static boolean isEncryptedPayload(String text) { + if (text == null || !Base64Util.isBase64(text)) { + return false; + } + try { + byte[] payload = Base64.getDecoder().decode(text); + return hasPayloadHeader(payload, AUTHENTICATED_PAYLOAD_HEADER) + || hasPayloadHeader(payload, CBC_PAYLOAD_HEADER); + } catch (Exception e) { + return false; + } + } + /** * Determine whether it is encrypted * @param text text diff --git a/hertzbeat-manager/src/main/java/org/apache/hertzbeat/manager/service/impl/ModelProviderConfigServiceImpl.java b/hertzbeat-manager/src/main/java/org/apache/hertzbeat/manager/service/impl/ModelProviderConfigServiceImpl.java index 0bbea9912a..1787d5b352 100644 --- a/hertzbeat-manager/src/main/java/org/apache/hertzbeat/manager/service/impl/ModelProviderConfigServiceImpl.java +++ b/hertzbeat-manager/src/main/java/org/apache/hertzbeat/manager/service/impl/ModelProviderConfigServiceImpl.java @@ -212,7 +212,7 @@ public class ModelProviderConfigServiceImpl extends AbstractGeneralConfigService } private String encryptSecret(String secret) { - if (!StringUtils.hasText(secret) || AesUtil.isCiphertext(secret)) { + if (!StringUtils.hasText(secret) || AesUtil.isCiphertext(secret) || AesUtil.isEncryptedPayload(secret)) { return secret; } String encrypted = AesUtil.aesEncode(secret); @@ -224,8 +224,11 @@ public class ModelProviderConfigServiceImpl extends AbstractGeneralConfigService private ModelProviderConfig decryptedConfiguration(ModelProviderConfig persisted) { ModelProviderConfig copy = copyConfiguration(persisted); - if (StringUtils.hasText(copy.getApiKey()) && AesUtil.isCiphertext(copy.getApiKey())) { + if (StringUtils.hasText(copy.getApiKey()) && AesUtil.isEncryptedPayload(copy.getApiKey())) { String ciphertext = copy.getApiKey(); + if (!AesUtil.isCiphertext(ciphertext)) { + throw new IllegalStateException("Model provider secret cannot be decrypted with the configured AES key"); + } String plaintext = AesUtil.aesDecode(ciphertext); if (Objects.equals(ciphertext, plaintext)) { throw new IllegalStateException("Model provider secret decryption failed"); diff --git a/hertzbeat-manager/src/test/java/org/apache/hertzbeat/manager/service/impl/ModelProviderConfigServiceImplTest.java b/hertzbeat-manager/src/test/java/org/apache/hertzbeat/manager/service/impl/ModelProviderConfigServiceImplTest.java index bc4e8d6ada..ba0929bd35 100644 --- a/hertzbeat-manager/src/test/java/org/apache/hertzbeat/manager/service/impl/ModelProviderConfigServiceImplTest.java +++ b/hertzbeat-manager/src/test/java/org/apache/hertzbeat/manager/service/impl/ModelProviderConfigServiceImplTest.java @@ -108,6 +108,22 @@ class ModelProviderConfigServiceImplTest { verify(generalConfigDao, never()).save(any(GeneralConfig.class)); } + @Test + void undecryptableCiphertextFailsInsteadOfLeakingToRuntime() { + AesUtil.setDefaultSecretKey("0123456789abcdef"); + ModelProviderConfigState created = service.createConfiguration(provider(API_KEY)); + String uid = created.getProviders().getFirst().getUid(); + String ciphertext = persistedState(stored.get().getContent()).getProviders().getFirst().getApiKey(); + + AesUtil.setDefaultSecretKey("fedcba9876543210"); + + assertThrows(IllegalStateException.class, () -> service.getConfiguration(uid)); + + service.updateConfiguration(uid, provider("")); + + assertEquals(ciphertext, persistedState(stored.get().getContent()).getProviders().getFirst().getApiKey()); + } + private ModelProviderConfig provider(String apiKey) { ModelProviderConfig config = new ModelProviderConfig(); config.setType("openai-compatible"); --------------------------------------------------------------------- To unsubscribe, e-mail: [email protected] For additional commands, e-mail: [email protected]
