From e83c9e3b0417c968fd6bb0bc81ca9168e03c0ae7 Mon Sep 17 00:00:00 2001 From: rithin-pullela-aws Date: Thu, 30 Apr 2026 14:49:42 -0700 Subject: [PATCH 1/3] feat: Add opt-in settings to enforce unique Agent and Agentic Memory Container names Adds two dynamic cluster settings, both defaulting to false for backward compatibility: - plugins.ml_commons.agent_name_uniqueness_enabled - plugins.ml_commons.agentic_memory_name_uniqueness_enabled When enabled, the respective registration/creation path rejects a request whose name collides with an existing resource in the same tenant, returning HTTP 409 CONFLICT. When disabled, behavior is unchanged. The uniqueness check queries the resource index by name.keyword via the remote-metadata SDK, so it honors multi-tenancy. IndexNotFoundException is treated as "no duplicate possible" to support the first-ever registration. Tenant validation is hoisted above the uniqueness check in TransportRegisterAgentAction#doExecute: in multi-tenant mode a missing tenantId now fails fast with 403 before any metadata search runs, rather than leaking existence via a cross-tenant 409. For PLAN_EXECUTE_AND_REFLECT agents registered without an existing executor_agent_id, the internally derived " (ReAct)" executor agent name is validated against the same tenant-aware search path, so the auto-created agent cannot silently collide with an existing name. The single-name check is extracted to checkAgentNameAvailable(name, tenantId) so both names share one implementation. Integration tests (RestMLAgentNameUniquenessIT, RestMLMemoryContainerNameUniquenessIT) cover end-to-end behavior against a running cluster for both flags: duplicates accepted when off, 409 when on, unique names always accepted, and dynamic flip without restart. These guard a regression class unit tests cannot catch (e.g. incorrect query field or malformed BoolQuery), since unit tests feed the transport action a synthetic SearchResponse rather than exercising the real search path. Signed-off-by: rithin-pullela-aws --- .../ml/common/settings/MLCommonsSettings.java | 20 ++ .../settings/MLFeatureEnabledSetting.java | 34 ++ .../MLFeatureEnabledSettingTests.java | 4 +- .../agents/TransportRegisterAgentAction.java | 112 ++++++- .../TransportCreateMemoryContainerAction.java | 83 ++++- .../ml/plugin/MachineLearningPlugin.java | 4 +- .../RegisterAgentTransportActionTests.java | 304 +++++++++++++++++- ...sportCreateMemoryContainerActionTests.java | 90 ++++++ .../ml/rest/RestMLAgentNameUniquenessIT.java | 114 +++++++ ...RestMLMemoryContainerNameUniquenessIT.java | 112 +++++++ 10 files changed, 866 insertions(+), 11 deletions(-) create mode 100644 plugin/src/test/java/org/opensearch/ml/rest/RestMLAgentNameUniquenessIT.java create mode 100644 plugin/src/test/java/org/opensearch/ml/rest/RestMLMemoryContainerNameUniquenessIT.java diff --git a/common/src/main/java/org/opensearch/ml/common/settings/MLCommonsSettings.java b/common/src/main/java/org/opensearch/ml/common/settings/MLCommonsSettings.java index 4540a5bdd2..45762b806e 100644 --- a/common/src/main/java/org/opensearch/ml/common/settings/MLCommonsSettings.java +++ b/common/src/main/java/org/opensearch/ml/common/settings/MLCommonsSettings.java @@ -539,4 +539,24 @@ private MLCommonsSettings() {} .boolSetting(ML_PLUGIN_SETTING_PREFIX + "ag_ui_enabled", false, Setting.Property.NodeScope, Setting.Property.Dynamic); public static final String ML_COMMONS_AG_UI_DISABLED_MESSAGE = "The AG-UI agent feature is not enabled. To enable, please update the setting " + ML_COMMONS_AG_UI_ENABLED.getKey(); + + // When enabled, registering an Agent with a name that already exists (per tenant) is rejected. + // Defaults to false for backward compatibility. + public static final Setting ML_COMMONS_AGENT_NAME_UNIQUENESS_ENABLED = Setting + .boolSetting( + ML_PLUGIN_SETTING_PREFIX + "agent_name_uniqueness_enabled", + false, + Setting.Property.NodeScope, + Setting.Property.Dynamic + ); + + // When enabled, creating an Agentic Memory Container with a name that already exists (per + // tenant) is rejected. Defaults to false for backward compatibility. + public static final Setting ML_COMMONS_AGENTIC_MEMORY_NAME_UNIQUENESS_ENABLED = Setting + .boolSetting( + ML_PLUGIN_SETTING_PREFIX + "agentic_memory_name_uniqueness_enabled", + false, + Setting.Property.NodeScope, + Setting.Property.Dynamic + ); } diff --git a/common/src/main/java/org/opensearch/ml/common/settings/MLFeatureEnabledSetting.java b/common/src/main/java/org/opensearch/ml/common/settings/MLFeatureEnabledSetting.java index cb7902a657..52df06d2f2 100644 --- a/common/src/main/java/org/opensearch/ml/common/settings/MLFeatureEnabledSetting.java +++ b/common/src/main/java/org/opensearch/ml/common/settings/MLFeatureEnabledSetting.java @@ -6,7 +6,9 @@ package org.opensearch.ml.common.settings; import static org.opensearch.ml.common.settings.MLCommonsSettings.ML_COMMONS_AGENTIC_MEMORY_ENABLED; +import static org.opensearch.ml.common.settings.MLCommonsSettings.ML_COMMONS_AGENTIC_MEMORY_NAME_UNIQUENESS_ENABLED; import static org.opensearch.ml.common.settings.MLCommonsSettings.ML_COMMONS_AGENT_FRAMEWORK_ENABLED; +import static org.opensearch.ml.common.settings.MLCommonsSettings.ML_COMMONS_AGENT_NAME_UNIQUENESS_ENABLED; import static org.opensearch.ml.common.settings.MLCommonsSettings.ML_COMMONS_AG_UI_ENABLED; import static org.opensearch.ml.common.settings.MLCommonsSettings.ML_COMMONS_CONNECTOR_PRIVATE_IP_ENABLED; import static org.opensearch.ml.common.settings.MLCommonsSettings.ML_COMMONS_CONTROLLER_ENABLED; @@ -76,6 +78,10 @@ public class MLFeatureEnabledSetting { private volatile Boolean isAGUIEnabled; + private volatile Boolean isAgentNameUniquenessEnabled; + + private volatile Boolean isAgenticMemoryNameUniquenessEnabled; + private final List listeners = new ArrayList<>(); public MLFeatureEnabledSetting(ClusterService clusterService, Settings settings) { @@ -101,6 +107,8 @@ public MLFeatureEnabledSetting(ClusterService clusterService, Settings settings) maxJsonSize = MLCommonsSettings.ML_COMMONS_MAX_JSON_SIZE.get(settings); isMcpHeaderPassthroughEnabled = ML_COMMONS_MCP_HEADER_PASSTHROUGH_ENABLED.get(settings); isAGUIEnabled = ML_COMMONS_AG_UI_ENABLED.get(settings); + isAgentNameUniquenessEnabled = ML_COMMONS_AGENT_NAME_UNIQUENESS_ENABLED.get(settings); + isAgenticMemoryNameUniquenessEnabled = ML_COMMONS_AGENTIC_MEMORY_NAME_UNIQUENESS_ENABLED.get(settings); clusterService .getClusterSettings() @@ -140,6 +148,12 @@ public MLFeatureEnabledSetting(ClusterService clusterService, Settings settings) .getClusterSettings() .addSettingsUpdateConsumer(ML_COMMONS_MCP_HEADER_PASSTHROUGH_ENABLED, it -> isMcpHeaderPassthroughEnabled = it); clusterService.getClusterSettings().addSettingsUpdateConsumer(ML_COMMONS_AG_UI_ENABLED, it -> isAGUIEnabled = it); + clusterService + .getClusterSettings() + .addSettingsUpdateConsumer(ML_COMMONS_AGENT_NAME_UNIQUENESS_ENABLED, it -> isAgentNameUniquenessEnabled = it); + clusterService + .getClusterSettings() + .addSettingsUpdateConsumer(ML_COMMONS_AGENTIC_MEMORY_NAME_UNIQUENESS_ENABLED, it -> isAgenticMemoryNameUniquenessEnabled = it); clusterService.getClusterSettings().addSettingsUpdateConsumer(ML_COMMONS_STATIC_METRIC_COLLECTION_ENABLED, it -> { isStaticMetricCollectionEnabled = it; for (SettingsChangeListener listener : listeners) { @@ -314,4 +328,24 @@ public boolean isMcpHeaderPassthroughEnabled() { public boolean isAGUIEnabled() { return isAGUIEnabled; } + + /** + * Whether agent name uniqueness is enforced. When enabled, registering an agent with a name + * that already exists within the same tenant is rejected. Defaults to false for backward + * compatibility. + * @return whether agent name uniqueness is enforced. + */ + public boolean isAgentNameUniquenessEnabled() { + return isAgentNameUniquenessEnabled; + } + + /** + * Whether agentic memory container name uniqueness is enforced. When enabled, creating a + * memory container with a name that already exists within the same tenant is rejected. + * Defaults to false for backward compatibility. + * @return whether agentic memory container name uniqueness is enforced. + */ + public boolean isAgenticMemoryNameUniquenessEnabled() { + return isAgenticMemoryNameUniquenessEnabled; + } } diff --git a/common/src/test/java/org/opensearch/ml/common/settings/MLFeatureEnabledSettingTests.java b/common/src/test/java/org/opensearch/ml/common/settings/MLFeatureEnabledSettingTests.java index ed811cac8c..0a26ff6d93 100644 --- a/common/src/test/java/org/opensearch/ml/common/settings/MLFeatureEnabledSettingTests.java +++ b/common/src/test/java/org/opensearch/ml/common/settings/MLFeatureEnabledSettingTests.java @@ -54,7 +54,9 @@ public void setUp() { MLCommonsSettings.ML_COMMONS_STREAM_ENABLED, MLCommonsSettings.ML_COMMONS_MAX_JSON_SIZE, MLCommonsSettings.ML_COMMONS_MCP_HEADER_PASSTHROUGH_ENABLED, - MLCommonsSettings.ML_COMMONS_AG_UI_ENABLED + MLCommonsSettings.ML_COMMONS_AG_UI_ENABLED, + MLCommonsSettings.ML_COMMONS_AGENT_NAME_UNIQUENESS_ENABLED, + MLCommonsSettings.ML_COMMONS_AGENTIC_MEMORY_NAME_UNIQUENESS_ENABLED ) ); when(mockClusterService.getClusterSettings()).thenReturn(mockClusterSettings); diff --git a/plugin/src/main/java/org/opensearch/ml/action/agents/TransportRegisterAgentAction.java b/plugin/src/main/java/org/opensearch/ml/action/agents/TransportRegisterAgentAction.java index de5a561ea2..e1e6282156 100644 --- a/plugin/src/main/java/org/opensearch/ml/action/agents/TransportRegisterAgentAction.java +++ b/plugin/src/main/java/org/opensearch/ml/action/agents/TransportRegisterAgentAction.java @@ -15,10 +15,12 @@ import java.util.HashMap; import java.util.Map; +import org.opensearch.ExceptionsHelper; import org.opensearch.OpenSearchException; import org.opensearch.OpenSearchStatusException; import org.opensearch.action.ActionRequest; import org.opensearch.action.index.IndexResponse; +import org.opensearch.action.search.SearchRequest; import org.opensearch.action.support.ActionFilters; import org.opensearch.action.support.HandledTransportAction; import org.opensearch.cluster.service.ClusterService; @@ -27,6 +29,9 @@ import org.opensearch.commons.authuser.User; import org.opensearch.core.action.ActionListener; import org.opensearch.core.rest.RestStatus; +import org.opensearch.index.IndexNotFoundException; +import org.opensearch.index.query.BoolQueryBuilder; +import org.opensearch.index.query.TermQueryBuilder; import org.opensearch.ml.action.agent.MLAgentRegistrationValidator; import org.opensearch.ml.action.contextmanagement.ContextManagementTemplateService; import org.opensearch.ml.common.MLAgentType; @@ -49,7 +54,9 @@ import org.opensearch.ml.utils.TenantAwareHelper; import org.opensearch.remote.metadata.client.PutDataObjectRequest; import org.opensearch.remote.metadata.client.SdkClient; +import org.opensearch.remote.metadata.client.SearchDataObjectRequest; import org.opensearch.remote.metadata.common.SdkClientUtils; +import org.opensearch.search.builder.SearchSourceBuilder; import org.opensearch.tasks.Task; import org.opensearch.transport.TransportService; import org.opensearch.transport.client.Client; @@ -99,13 +106,107 @@ protected void doExecute(Task task, ActionRequest request, ActionListener { + // Check if this agent needs model creation + if (mlAgent.usesUnifiedInterface()) { + createModelAndRegisterAgent(mlAgent, listener); + return; + } + registerAgent(mlAgent, listener); + }, listener::onFailure)); + } + + /** + * When {@code plugins.ml_commons.agent_name_uniqueness_enabled} is enabled, reject the + * registration if an agent with the same name already exists in the same tenant. When the + * setting is disabled (default), this is a no-op so existing clusters remain backward compatible. + * + *

For PLAN_EXECUTE_AND_REFLECT agents without a pre-existing executor-agent-id, an internal + * "{@code (ReAct)}" executor agent is auto-created; this method also checks that + * derived name so the auto-created agent cannot collide with an existing one. + * + *

Note: this is a best-effort check, not a transactional guard. Two concurrent register + * requests with the same name can both pass this check before either write is visible. + * See the PR description for a follow-up plan if stricter semantics are required. + */ + private void validateAgentNameUniqueness(MLAgent mlAgent, ActionListener listener) { + if (!mlFeatureEnabledSetting.isAgentNameUniquenessEnabled()) { + listener.onResponse(null); + return; + } + + // MLAgent.validate() already rejects null/blank/over-length names upstream, so we rely on + // that invariant here and don't re-check. + String name = mlAgent.getName(); + String tenantId = mlAgent.getTenantId(); + checkAgentNameAvailable(name, tenantId, ActionListener.wrap(unused -> { + // If this is a PLAN_EXECUTE_AND_REFLECT registration that will auto-create an + // executor agent named " (ReAct)", validate that derived name too. + if (MLAgentType.from(mlAgent.getType()) == MLAgentType.PLAN_EXECUTE_AND_REFLECT + && mlAgent.getParameters() != null + && !mlAgent.getParameters().containsKey(MLPlanExecuteAndReflectAgentRunner.EXECUTOR_AGENT_ID_FIELD)) { + checkAgentNameAvailable(name + " (ReAct)", tenantId, listener); + } else { + listener.onResponse(null); + } + }, listener::onFailure)); + } + + private void checkAgentNameAvailable(String name, String tenantId, ActionListener listener) { + BoolQueryBuilder query = new BoolQueryBuilder().filter(new TermQueryBuilder("name.keyword", name)); + SearchSourceBuilder sourceBuilder = new SearchSourceBuilder().query(query).size(1).fetchSource(false); + SearchRequest searchRequest = new SearchRequest(ML_AGENT_INDEX).source(sourceBuilder); + SearchDataObjectRequest searchDataObjectRequest = SearchDataObjectRequest + .builder() + .indices(searchRequest.indices()) + .searchSourceBuilder(searchRequest.source()) + .tenantId(tenantId) + .build(); + + try (ThreadContext.StoredContext context = client.threadPool().getThreadContext().stashContext()) { + sdkClient.searchDataObjectAsync(searchDataObjectRequest).whenComplete((r, throwable) -> { + context.restore(); + if (throwable != null) { + if (ExceptionsHelper.unwrap(throwable, IndexNotFoundException.class) != null) { + // Index not yet created - no duplicate possible + listener.onResponse(null); + return; + } + Exception cause = SdkClientUtils.unwrapAndConvertToException(throwable); + log.error("Failed to search ML agent index for name uniqueness check", cause); + listener.onFailure(cause); + return; + } + try { + long totalHits = r.searchResponse().getHits().getTotalHits() == null + ? 0 + : r.searchResponse().getHits().getTotalHits().value(); + if (totalHits > 0) { + listener + .onFailure( + new OpenSearchStatusException( + "An agent with name [" + name + "] already exists. Agent names must be unique.", + RestStatus.CONFLICT + ) + ); + } else { + listener.onResponse(null); + } + } catch (Exception e) { + log.error("Failed to parse search response for agent name uniqueness check", e); + listener.onFailure(e); + } + }); + } catch (Exception e) { + log.error("Failed to execute agent name uniqueness check", e); + listener.onFailure(e); + } } private void createModelAndRegisterAgent(MLAgent mlAgent, ActionListener listener) { @@ -185,9 +286,6 @@ private void proceedWithAgentRegistration(MLAgent agent, ActionListener { + // Validate configuration before creating memory container + validateConfigurationAndCreate(input, user, tenantId, listener); + }, listener::onFailure)); + } + + private void validateConfigurationAndCreate( + MLCreateMemoryContainerInput input, + User user, + String tenantId, + ActionListener listener + ) { validateConfiguration(input.getConfiguration(), ActionListener.wrap(isValid -> { // Check if memory container index exists, create if not ActionListener indexCheckListener = ActionListener.wrap(created -> { @@ -288,6 +306,69 @@ private void indexMemoryContainer(MLMemoryContainer container, ActionListener listener) { + if (!mlFeatureEnabledSetting.isAgenticMemoryNameUniquenessEnabled()) { + listener.onResponse(null); + return; + } + + // MLCreateMemoryContainerInput rejects null names upstream, so we rely on that invariant + // here and don't re-check. + BoolQueryBuilder query = new BoolQueryBuilder().filter(new TermQueryBuilder("name.keyword", name)); + SearchSourceBuilder sourceBuilder = new SearchSourceBuilder().query(query).size(1).fetchSource(false); + SearchRequest searchRequest = new SearchRequest(ML_MEMORY_CONTAINER_INDEX).source(sourceBuilder); + SearchDataObjectRequest searchDataObjectRequest = SearchDataObjectRequest + .builder() + .indices(searchRequest.indices()) + .searchSourceBuilder(searchRequest.source()) + .tenantId(tenantId) + .build(); + + try (ThreadContext.StoredContext context = client.threadPool().getThreadContext().stashContext()) { + sdkClient.searchDataObjectAsync(searchDataObjectRequest).whenComplete((r, throwable) -> { + context.restore(); + if (throwable != null) { + if (ExceptionsHelper.unwrap(throwable, IndexNotFoundException.class) != null) { + // Index not yet created - no duplicate possible + listener.onResponse(null); + return; + } + Exception cause = SdkClientUtils.unwrapAndConvertToException(throwable); + log.error("Failed to search memory container index for name uniqueness check", cause); + listener.onFailure(cause); + return; + } + try { + long totalHits = r.searchResponse().getHits().getTotalHits() == null + ? 0 + : r.searchResponse().getHits().getTotalHits().value(); + if (totalHits > 0) { + listener + .onFailure( + new OpenSearchStatusException( + "A memory container with name [" + name + "] already exists. Memory container names must be unique.", + RestStatus.CONFLICT + ) + ); + } else { + listener.onResponse(null); + } + } catch (Exception e) { + log.error("Failed to parse search response for memory container name uniqueness check", e); + listener.onFailure(e); + } + }); + } catch (Exception e) { + log.error("Failed to execute memory container name uniqueness check", e); + listener.onFailure(e); + } + } + private void validateConfiguration(MemoryConfiguration config, ActionListener listener) { // Validate that strategies have required AI models try { diff --git a/plugin/src/main/java/org/opensearch/ml/plugin/MachineLearningPlugin.java b/plugin/src/main/java/org/opensearch/ml/plugin/MachineLearningPlugin.java index 60004d8834..06beeed11d 100644 --- a/plugin/src/main/java/org/opensearch/ml/plugin/MachineLearningPlugin.java +++ b/plugin/src/main/java/org/opensearch/ml/plugin/MachineLearningPlugin.java @@ -1434,7 +1434,9 @@ public List> getSettings() { MLCommonsSettings.ML_COMMONS_MAX_JSON_SIZE, MLCommonsSettings.ML_COMMONS_UNIFIED_AGENT_API_ENABLED, MLCommonsSettings.ML_COMMONS_MCP_HEADER_PASSTHROUGH_ENABLED, - MLCommonsSettings.ML_COMMONS_AG_UI_ENABLED + MLCommonsSettings.ML_COMMONS_AG_UI_ENABLED, + MLCommonsSettings.ML_COMMONS_AGENT_NAME_UNIQUENESS_ENABLED, + MLCommonsSettings.ML_COMMONS_AGENTIC_MEMORY_NAME_UNIQUENESS_ENABLED ); return settings; } diff --git a/plugin/src/test/java/org/opensearch/ml/action/agents/RegisterAgentTransportActionTests.java b/plugin/src/test/java/org/opensearch/ml/action/agents/RegisterAgentTransportActionTests.java index 13972e8cac..5df006b3a1 100644 --- a/plugin/src/test/java/org/opensearch/ml/action/agents/RegisterAgentTransportActionTests.java +++ b/plugin/src/test/java/org/opensearch/ml/action/agents/RegisterAgentTransportActionTests.java @@ -14,6 +14,7 @@ import static org.mockito.Mockito.when; import static org.opensearch.ml.common.CommonValue.MCP_CONNECTORS_FIELD; import static org.opensearch.ml.common.CommonValue.ML_AGENT_INDEX; +import static org.opensearch.ml.common.settings.MLCommonsSettings.ML_COMMONS_AGENT_NAME_UNIQUENESS_ENABLED; import static org.opensearch.ml.common.settings.MLCommonsSettings.ML_COMMONS_MCP_CONNECTOR_DISABLED_MESSAGE; import static org.opensearch.ml.common.settings.MLCommonsSettings.ML_COMMONS_MCP_CONNECTOR_ENABLED; import static org.opensearch.ml.engine.algorithms.agent.MLChatAgentRunner.LLM_INTERFACE; @@ -111,7 +112,8 @@ public void setup() throws IOException { when(client.threadPool()).thenReturn(threadPool); when(threadPool.getThreadContext()).thenReturn(threadContext); when(clusterService.getSettings()).thenReturn(settings); - when(this.clusterService.getClusterSettings()).thenReturn(new ClusterSettings(settings, Set.of(ML_COMMONS_MCP_CONNECTOR_ENABLED))); + when(this.clusterService.getClusterSettings()) + .thenReturn(new ClusterSettings(settings, Set.of(ML_COMMONS_MCP_CONNECTOR_ENABLED, ML_COMMONS_AGENT_NAME_UNIQUENESS_ENABLED))); transportRegisterAgentAction = new TransportRegisterAgentAction( transportService, actionFilters, @@ -609,4 +611,304 @@ public void test_execute_registerAgent_WithModelSpec_ModelRegistrationFailure() verify(actionListener).onFailure(argumentCaptor.capture()); assertEquals("Model registration failed", argumentCaptor.getValue().getMessage()); } + + @Test + public void test_execute_registerAgent_uniquenessEnforced_duplicateNameRejected() { + when(mlFeatureEnabledSetting.isAgentNameUniquenessEnabled()).thenReturn(true); + + MLRegisterAgentRequest request = mock(MLRegisterAgentRequest.class); + MLAgent mlAgent = MLAgent + .builder() + .name("duplicate-agent") + .type(MLAgentType.CONVERSATIONAL.name()) + .description("description") + .llm(new LLMSpec("model_id", new HashMap<>())) + .build(); + when(request.getMlAgent()).thenReturn(mlAgent); + + // Simulate a search hit (name already exists) + doAnswer(invocation -> { + ActionListener al = invocation.getArgument(1); + org.apache.lucene.search.TotalHits totalHits = new org.apache.lucene.search.TotalHits( + 1L, + org.apache.lucene.search.TotalHits.Relation.EQUAL_TO + ); + org.opensearch.search.SearchHits hits = new org.opensearch.search.SearchHits( + new org.opensearch.search.SearchHit[0], + totalHits, + Float.NaN + ); + org.opensearch.search.internal.InternalSearchResponse internal = new org.opensearch.search.internal.InternalSearchResponse( + hits, + org.opensearch.search.aggregations.InternalAggregations.EMPTY, + null, + null, + false, + null, + 0 + ); + org.opensearch.action.search.SearchResponse searchResponse = new org.opensearch.action.search.SearchResponse( + internal, + null, + 1, + 1, + 0, + 1, + org.opensearch.action.search.ShardSearchFailure.EMPTY_ARRAY, + org.opensearch.action.search.SearchResponse.Clusters.EMPTY + ); + al.onResponse(searchResponse); + return null; + }).when(client).search(any(), any()); + + transportRegisterAgentAction.doExecute(task, request, actionListener); + + ArgumentCaptor argumentCaptor = ArgumentCaptor.forClass(Exception.class); + verify(actionListener).onFailure(argumentCaptor.capture()); + assertTrue(argumentCaptor.getValue().getMessage().contains("already exists")); + // Index creation must NOT have been attempted when a duplicate is found + verify(mlIndicesHandler, times(0)).initMLAgentIndex(any()); + } + + @Test + public void test_execute_registerAgent_uniquenessEnforced_uniqueNameAllowed() { + when(mlFeatureEnabledSetting.isAgentNameUniquenessEnabled()).thenReturn(true); + + MLRegisterAgentRequest request = mock(MLRegisterAgentRequest.class); + MLAgent mlAgent = MLAgent + .builder() + .name("unique-agent") + .type(MLAgentType.CONVERSATIONAL.name()) + .description("description") + .llm(new LLMSpec("model_id", new HashMap<>())) + .build(); + when(request.getMlAgent()).thenReturn(mlAgent); + + // Simulate a search returning zero hits (name is unique) + doAnswer(invocation -> { + ActionListener al = invocation.getArgument(1); + org.apache.lucene.search.TotalHits totalHits = new org.apache.lucene.search.TotalHits( + 0L, + org.apache.lucene.search.TotalHits.Relation.EQUAL_TO + ); + org.opensearch.search.SearchHits hits = new org.opensearch.search.SearchHits( + new org.opensearch.search.SearchHit[0], + totalHits, + Float.NaN + ); + org.opensearch.search.internal.InternalSearchResponse internal = new org.opensearch.search.internal.InternalSearchResponse( + hits, + org.opensearch.search.aggregations.InternalAggregations.EMPTY, + null, + null, + false, + null, + 0 + ); + org.opensearch.action.search.SearchResponse searchResponse = new org.opensearch.action.search.SearchResponse( + internal, + null, + 1, + 1, + 0, + 1, + org.opensearch.action.search.ShardSearchFailure.EMPTY_ARRAY, + org.opensearch.action.search.SearchResponse.Clusters.EMPTY + ); + al.onResponse(searchResponse); + return null; + }).when(client).search(any(), any()); + + doAnswer(invocation -> { + ActionListener listener = invocation.getArgument(0); + listener.onResponse(true); + return null; + }).when(mlIndicesHandler).initMLAgentIndex(any()); + + doAnswer(invocation -> { + ActionListener al = invocation.getArgument(1); + al.onResponse(indexResponse); + return null; + }).when(client).index(any(), any()); + + transportRegisterAgentAction.doExecute(task, request, actionListener); + + ArgumentCaptor argumentCaptor = ArgumentCaptor.forClass(MLRegisterAgentResponse.class); + verify(actionListener).onResponse(argumentCaptor.capture()); + assertNotNull(argumentCaptor.getValue()); + } + + @Test + public void test_execute_registerAgent_uniquenessEnforced_indexNotFound_allowsRegistration() { + when(mlFeatureEnabledSetting.isAgentNameUniquenessEnabled()).thenReturn(true); + + MLRegisterAgentRequest request = mock(MLRegisterAgentRequest.class); + MLAgent mlAgent = MLAgent + .builder() + .name("first-agent") + .type(MLAgentType.CONVERSATIONAL.name()) + .description("description") + .llm(new LLMSpec("model_id", new HashMap<>())) + .build(); + when(request.getMlAgent()).thenReturn(mlAgent); + + // Simulate IndexNotFoundException from search - happens before any agent has ever been registered. + doAnswer(invocation -> { + ActionListener al = invocation.getArgument(1); + al.onFailure(new org.opensearch.index.IndexNotFoundException(ML_AGENT_INDEX)); + return null; + }).when(client).search(any(), any()); + + doAnswer(invocation -> { + ActionListener listener = invocation.getArgument(0); + listener.onResponse(true); + return null; + }).when(mlIndicesHandler).initMLAgentIndex(any()); + + doAnswer(invocation -> { + ActionListener al = invocation.getArgument(1); + al.onResponse(indexResponse); + return null; + }).when(client).index(any(), any()); + + transportRegisterAgentAction.doExecute(task, request, actionListener); + + ArgumentCaptor argumentCaptor = ArgumentCaptor.forClass(MLRegisterAgentResponse.class); + verify(actionListener).onResponse(argumentCaptor.capture()); + assertNotNull(argumentCaptor.getValue()); + } + + @Test + public void test_execute_registerAgent_uniquenessDisabled_skipsSearch() { + // Default: uniqueness disabled - no search should occur + when(mlFeatureEnabledSetting.isAgentNameUniquenessEnabled()).thenReturn(false); + + MLRegisterAgentRequest request = mock(MLRegisterAgentRequest.class); + MLAgent mlAgent = MLAgent + .builder() + .name("any-agent") + .type(MLAgentType.CONVERSATIONAL.name()) + .description("description") + .llm(new LLMSpec("model_id", new HashMap<>())) + .build(); + when(request.getMlAgent()).thenReturn(mlAgent); + + doAnswer(invocation -> { + ActionListener listener = invocation.getArgument(0); + listener.onResponse(true); + return null; + }).when(mlIndicesHandler).initMLAgentIndex(any()); + + doAnswer(invocation -> { + ActionListener al = invocation.getArgument(1); + al.onResponse(indexResponse); + return null; + }).when(client).index(any(), any()); + + transportRegisterAgentAction.doExecute(task, request, actionListener); + + verify(client, times(0)).search(any(), any()); + ArgumentCaptor argumentCaptor = ArgumentCaptor.forClass(MLRegisterAgentResponse.class); + verify(actionListener).onResponse(argumentCaptor.capture()); + } + + /** + * Builds a real SearchResponse with the given total-hit count so we can drive the + * ActionListener callback path through mock client.search(...). + */ + private org.opensearch.action.search.SearchResponse buildSearchResponseWithHits(long hitCount) { + org.apache.lucene.search.TotalHits totalHits = new org.apache.lucene.search.TotalHits( + hitCount, + org.apache.lucene.search.TotalHits.Relation.EQUAL_TO + ); + org.opensearch.search.SearchHits hits = new org.opensearch.search.SearchHits( + new org.opensearch.search.SearchHit[0], + totalHits, + Float.NaN + ); + org.opensearch.search.internal.InternalSearchResponse internal = new org.opensearch.search.internal.InternalSearchResponse( + hits, + org.opensearch.search.aggregations.InternalAggregations.EMPTY, + null, + null, + false, + null, + 0 + ); + return new org.opensearch.action.search.SearchResponse( + internal, + null, + 1, + 1, + 0, + 1, + org.opensearch.action.search.ShardSearchFailure.EMPTY_ARRAY, + org.opensearch.action.search.SearchResponse.Clusters.EMPTY + ); + } + + @Test + public void test_execute_registerAgent_uniquenessEnforced_planExecuteAndReflect_executorNameCollision_rejected() { + when(mlFeatureEnabledSetting.isAgentNameUniquenessEnabled()).thenReturn(true); + + Map parameters = new HashMap<>(); + parameters.put("tools", "[]"); + MLAgent mlAgent = MLAgent + .builder() + .name("per_agent") + .type(MLAgentType.PLAN_EXECUTE_AND_REFLECT.name()) + .description("Plan-execute-and-reflect agent") + .parameters(parameters) + .llm(new LLMSpec("test-model-id", new HashMap<>())) + .build(); + + MLRegisterAgentRequest request = mock(MLRegisterAgentRequest.class); + when(request.getMlAgent()).thenReturn(mlAgent); + + // First search: submitted name is available (0 hits). Second search: derived executor + // name "per_agent (ReAct)" already exists (1 hit) -> request must fail before any index write. + doAnswer(new org.mockito.stubbing.Answer() { + private int callCount = 0; + + @Override + public Void answer(org.mockito.invocation.InvocationOnMock invocation) { + ActionListener al = invocation.getArgument(1); + long hits = (callCount++ == 0) ? 0L : 1L; + al.onResponse(buildSearchResponseWithHits(hits)); + return null; + } + }).when(client).search(any(), any()); + + transportRegisterAgentAction.doExecute(task, request, actionListener); + + ArgumentCaptor argumentCaptor = ArgumentCaptor.forClass(Exception.class); + verify(actionListener).onFailure(argumentCaptor.capture()); + assertTrue(argumentCaptor.getValue().getMessage().contains("(ReAct)")); + assertTrue(argumentCaptor.getValue().getMessage().contains("already exists")); + // Neither the submitted agent nor its executor should have been indexed + verify(mlIndicesHandler, times(0)).initMLAgentIndex(any()); + } + + @Test + public void test_execute_registerAgent_multiTenancy_missingTenantId_failsBeforeSearch() { + when(mlFeatureEnabledSetting.isMultiTenancyEnabled()).thenReturn(true); + when(mlFeatureEnabledSetting.isAgentNameUniquenessEnabled()).thenReturn(true); + + // No tenantId on the agent + MLAgent mlAgent = MLAgent + .builder() + .name("some-agent") + .type(MLAgentType.CONVERSATIONAL.name()) + .description("description") + .llm(new LLMSpec("model_id", new HashMap<>())) + .build(); + MLRegisterAgentRequest request = mock(MLRegisterAgentRequest.class); + when(request.getMlAgent()).thenReturn(mlAgent); + + transportRegisterAgentAction.doExecute(task, request, actionListener); + + // Tenant validation must fail fast, before any metadata search is issued + verify(client, times(0)).search(any(), any()); + verify(actionListener).onFailure(any(Exception.class)); + } } diff --git a/plugin/src/test/java/org/opensearch/ml/action/memorycontainer/TransportCreateMemoryContainerActionTests.java b/plugin/src/test/java/org/opensearch/ml/action/memorycontainer/TransportCreateMemoryContainerActionTests.java index 118284b5b9..2468bdc5db 100644 --- a/plugin/src/test/java/org/opensearch/ml/action/memorycontainer/TransportCreateMemoryContainerActionTests.java +++ b/plugin/src/test/java/org/opensearch/ml/action/memorycontainer/TransportCreateMemoryContainerActionTests.java @@ -75,6 +75,8 @@ import org.opensearch.remote.metadata.client.PutDataObjectRequest; import org.opensearch.remote.metadata.client.PutDataObjectResponse; import org.opensearch.remote.metadata.client.SdkClient; +import org.opensearch.remote.metadata.client.SearchDataObjectRequest; +import org.opensearch.remote.metadata.client.SearchDataObjectResponse; import org.opensearch.tasks.Task; import org.opensearch.test.OpenSearchTestCase; import org.opensearch.threadpool.ThreadPool; @@ -1956,4 +1958,92 @@ private void mockSuccessfulLLMValidation() { return null; }).when(mlModelManager).getModel(eq("test-embedding-model"), any()); } + + private SearchDataObjectResponse searchDataObjectResponseWithHits(long hitCount) { + org.apache.lucene.search.TotalHits totalHits = new org.apache.lucene.search.TotalHits( + hitCount, + org.apache.lucene.search.TotalHits.Relation.EQUAL_TO + ); + org.opensearch.search.SearchHits hits = new org.opensearch.search.SearchHits( + new org.opensearch.search.SearchHit[0], + totalHits, + Float.NaN + ); + org.opensearch.search.internal.InternalSearchResponse internal = new org.opensearch.search.internal.InternalSearchResponse( + hits, + org.opensearch.search.aggregations.InternalAggregations.EMPTY, + null, + null, + false, + null, + 0 + ); + org.opensearch.action.search.SearchResponse searchResponse = new org.opensearch.action.search.SearchResponse( + internal, + null, + 1, + 1, + 0, + 1, + org.opensearch.action.search.ShardSearchFailure.EMPTY_ARRAY, + org.opensearch.action.search.SearchResponse.Clusters.EMPTY + ); + SearchDataObjectResponse sdkResp = mock(SearchDataObjectResponse.class); + when(sdkResp.searchResponse()).thenReturn(searchResponse); + return sdkResp; + } + + public void testDoExecute_UniquenessEnforced_DuplicateNameRejected() throws InterruptedException { + when(mlFeatureEnabledSetting.isAgenticMemoryNameUniquenessEnabled()).thenReturn(true); + + // Simulate search returning 1 hit (name already exists) + CompletableFuture future = CompletableFuture.completedFuture(searchDataObjectResponseWithHits(1L)); + when(sdkClient.searchDataObjectAsync(any(SearchDataObjectRequest.class))).thenReturn(future); + + action.doExecute(task, request, actionListener); + + verify(actionListener).onFailure(exceptionCaptor.capture()); + Exception exception = exceptionCaptor.getValue(); + assertNotNull(exception); + assertTrue(exception instanceof OpenSearchStatusException); + assertEquals(RestStatus.CONFLICT, ((OpenSearchStatusException) exception).status()); + assertTrue(exception.getMessage().contains("already exists")); + // putDataObjectAsync must NOT be called when duplicate is detected + verify(sdkClient, org.mockito.Mockito.never()).putDataObjectAsync(any(PutDataObjectRequest.class)); + } + + public void testDoExecute_UniquenessEnforced_UniqueNameAllowed() throws InterruptedException { + when(mlFeatureEnabledSetting.isAgenticMemoryNameUniquenessEnabled()).thenReturn(true); + + // Simulate search returning 0 hits (name is unique) + CompletableFuture searchFuture = CompletableFuture.completedFuture(searchDataObjectResponseWithHits(0L)); + when(sdkClient.searchDataObjectAsync(any(SearchDataObjectRequest.class))).thenReturn(searchFuture); + + mockSuccessfulCreatePipeline(); + mockAndRunExecuteMethod(request); + + verify(sdkClient).searchDataObjectAsync(any(SearchDataObjectRequest.class)); + } + + public void testDoExecute_UniquenessEnforced_IndexNotFound_AllowsCreation() throws InterruptedException { + when(mlFeatureEnabledSetting.isAgenticMemoryNameUniquenessEnabled()).thenReturn(true); + + // Simulate IndexNotFoundException - first time a memory container is ever created + CompletableFuture searchFuture = new CompletableFuture<>(); + searchFuture.completeExceptionally(new IndexNotFoundException(ML_MEMORY_CONTAINER_INDEX)); + when(sdkClient.searchDataObjectAsync(any(SearchDataObjectRequest.class))).thenReturn(searchFuture); + + mockSuccessfulCreatePipeline(); + mockAndRunExecuteMethod(request); + } + + public void testDoExecute_UniquenessDisabled_SkipsSearch() throws InterruptedException { + // Default: uniqueness disabled + when(mlFeatureEnabledSetting.isAgenticMemoryNameUniquenessEnabled()).thenReturn(false); + + mockSuccessfulCreatePipeline(); + mockAndRunExecuteMethod(request); + + verify(sdkClient, org.mockito.Mockito.never()).searchDataObjectAsync(any(SearchDataObjectRequest.class)); + } } diff --git a/plugin/src/test/java/org/opensearch/ml/rest/RestMLAgentNameUniquenessIT.java b/plugin/src/test/java/org/opensearch/ml/rest/RestMLAgentNameUniquenessIT.java new file mode 100644 index 0000000000..64fab72d75 --- /dev/null +++ b/plugin/src/test/java/org/opensearch/ml/rest/RestMLAgentNameUniquenessIT.java @@ -0,0 +1,114 @@ +/* + * Copyright OpenSearch Contributors + * SPDX-License-Identifier: Apache-2.0 + */ + +package org.opensearch.ml.rest; + +import java.io.IOException; + +import org.junit.After; +import org.junit.Before; +import org.opensearch.client.Response; +import org.opensearch.client.ResponseException; +import org.opensearch.ml.utils.TestHelper; + +/** + * End-to-end integration tests for the agent-name-uniqueness feature gated by + * {@code plugins.ml_commons.agent_name_uniqueness_enabled}. + */ +public class RestMLAgentNameUniquenessIT extends MLCommonsRestTestCase { + + private static final String SETTING_KEY = "plugins.ml_commons.agent_name_uniqueness_enabled"; + + @Before + public void setupUniquenessSetting() throws IOException { + // Start every test with the flag explicitly off so tests are independent of persisted state. + updateClusterSettings(SETTING_KEY, false); + } + + @After + public void resetUniquenessSetting() throws IOException { + updateClusterSettings(SETTING_KEY, false); + } + + public void testDuplicateAllowed_WhenFlagOff() throws IOException { + String body = registerAgentBody("it-agent-flag-off"); + + Response r1 = registerAgent(body); + assertEquals(200, r1.getStatusLine().getStatusCode()); + + Response r2 = registerAgent(body); + assertEquals(200, r2.getStatusLine().getStatusCode()); + + String id1 = (String) parseResponseToMap(r1).get("agent_id"); + String id2 = (String) parseResponseToMap(r2).get("agent_id"); + assertNotNull(id1); + assertNotNull(id2); + assertNotEquals("Duplicate names with flag off should still produce distinct agent IDs", id1, id2); + } + + public void testDuplicateRejected_WhenFlagOn() throws IOException { + updateClusterSettings(SETTING_KEY, true); + + String body = registerAgentBody("it-agent-flag-on"); + + Response r1 = registerAgent(body); + assertEquals(200, r1.getStatusLine().getStatusCode()); + + try { + registerAgent(body); + fail("Expected duplicate registration to be rejected with 409 when uniqueness flag is on"); + } catch (ResponseException e) { + assertEquals(409, e.getResponse().getStatusLine().getStatusCode()); + String payload = TestHelper.httpEntityToString(e.getResponse().getEntity()); + assertTrue( + "409 response should cite the duplicate name, got: " + payload, + payload.contains("already exists") && payload.contains("it-agent-flag-on") + ); + } + } + + public void testUniqueNameAccepted_WhenFlagOn() throws IOException { + updateClusterSettings(SETTING_KEY, true); + + Response r = registerAgent(registerAgentBody("it-agent-unique-" + System.nanoTime())); + assertEquals(200, r.getStatusLine().getStatusCode()); + assertNotNull(parseResponseToMap(r).get("agent_id")); + } + + public void testDynamicFlip_OnThenOff() throws IOException { + String body = registerAgentBody("it-agent-dynamic-flip"); + + updateClusterSettings(SETTING_KEY, true); + Response r1 = registerAgent(body); + assertEquals(200, r1.getStatusLine().getStatusCode()); + + try { + registerAgent(body); + fail("Expected 409 while flag is on"); + } catch (ResponseException e) { + assertEquals(409, e.getResponse().getStatusLine().getStatusCode()); + } + + // Flip off; duplicate should now be accepted again without restart. + updateClusterSettings(SETTING_KEY, false); + Response r2 = registerAgent(body); + assertEquals("Duplicate must be accepted again once flag is flipped off", 200, r2.getStatusLine().getStatusCode()); + } + + private static Response registerAgent(String body) throws IOException { + return TestHelper.makeRequest(client(), "POST", "/_plugins/_ml/agents/_register", null, TestHelper.toHttpEntity(body), null); + } + + private static String registerAgentBody(String name) { + return "{\n" + + " \"name\": \"" + + name + + "\",\n" + + " \"type\": \"flow\",\n" + + " \"description\": \"uniqueness IT agent\",\n" + + " \"tools\": []\n" + + "}"; + } +} diff --git a/plugin/src/test/java/org/opensearch/ml/rest/RestMLMemoryContainerNameUniquenessIT.java b/plugin/src/test/java/org/opensearch/ml/rest/RestMLMemoryContainerNameUniquenessIT.java new file mode 100644 index 0000000000..485d94bae4 --- /dev/null +++ b/plugin/src/test/java/org/opensearch/ml/rest/RestMLMemoryContainerNameUniquenessIT.java @@ -0,0 +1,112 @@ +/* + * Copyright OpenSearch Contributors + * SPDX-License-Identifier: Apache-2.0 + */ + +package org.opensearch.ml.rest; + +import java.io.IOException; + +import org.junit.After; +import org.junit.Before; +import org.opensearch.client.Response; +import org.opensearch.client.ResponseException; +import org.opensearch.ml.utils.TestHelper; + +/** + * End-to-end integration tests for the memory-container-name-uniqueness feature gated by + * {@code plugins.ml_commons.agentic_memory_name_uniqueness_enabled}. + */ +public class RestMLMemoryContainerNameUniquenessIT extends MLCommonsRestTestCase { + + private static final String SETTING_KEY = "plugins.ml_commons.agentic_memory_name_uniqueness_enabled"; + + @Before + public void setupUniquenessSetting() throws IOException { + updateClusterSettings(SETTING_KEY, false); + } + + @After + public void resetUniquenessSetting() throws IOException { + updateClusterSettings(SETTING_KEY, false); + } + + public void testDuplicateAllowed_WhenFlagOff() throws IOException { + String body = createMemoryContainerBody("it-memcontainer-flag-off"); + + Response r1 = createMemoryContainer(body); + assertEquals(200, r1.getStatusLine().getStatusCode()); + + Response r2 = createMemoryContainer(body); + assertEquals(200, r2.getStatusLine().getStatusCode()); + + String id1 = (String) parseResponseToMap(r1).get("memory_container_id"); + String id2 = (String) parseResponseToMap(r2).get("memory_container_id"); + assertNotNull(id1); + assertNotNull(id2); + assertNotEquals("Duplicate names with flag off should still produce distinct container IDs", id1, id2); + } + + public void testDuplicateRejected_WhenFlagOn() throws IOException { + updateClusterSettings(SETTING_KEY, true); + + String body = createMemoryContainerBody("it-memcontainer-flag-on"); + + Response r1 = createMemoryContainer(body); + assertEquals(200, r1.getStatusLine().getStatusCode()); + + try { + createMemoryContainer(body); + fail("Expected duplicate memory container creation to be rejected with 409 when uniqueness flag is on"); + } catch (ResponseException e) { + assertEquals(409, e.getResponse().getStatusLine().getStatusCode()); + String payload = TestHelper.httpEntityToString(e.getResponse().getEntity()); + assertTrue( + "409 response should cite the duplicate name, got: " + payload, + payload.contains("already exists") && payload.contains("it-memcontainer-flag-on") + ); + } + } + + public void testUniqueNameAccepted_WhenFlagOn() throws IOException { + updateClusterSettings(SETTING_KEY, true); + + Response r = createMemoryContainer(createMemoryContainerBody("it-memcontainer-unique-" + System.nanoTime())); + assertEquals(200, r.getStatusLine().getStatusCode()); + assertNotNull(parseResponseToMap(r).get("memory_container_id")); + } + + public void testDynamicFlip_OnThenOff() throws IOException { + String body = createMemoryContainerBody("it-memcontainer-dynamic-flip"); + + updateClusterSettings(SETTING_KEY, true); + Response r1 = createMemoryContainer(body); + assertEquals(200, r1.getStatusLine().getStatusCode()); + + try { + createMemoryContainer(body); + fail("Expected 409 while flag is on"); + } catch (ResponseException e) { + assertEquals(409, e.getResponse().getStatusLine().getStatusCode()); + } + + updateClusterSettings(SETTING_KEY, false); + Response r2 = createMemoryContainer(body); + assertEquals("Duplicate must be accepted again once flag is flipped off", 200, r2.getStatusLine().getStatusCode()); + } + + private static Response createMemoryContainer(String body) throws IOException { + return TestHelper + .makeRequest(client(), "POST", "/_plugins/_ml/memory_containers/_create", null, TestHelper.toHttpEntity(body), null); + } + + private static String createMemoryContainerBody(String name) { + return "{\n" + + " \"name\": \"" + + name + + "\",\n" + + " \"description\": \"uniqueness IT memory container\",\n" + + " \"configuration\": {}\n" + + "}"; + } +} From d48f3c6a0643ac52f9caea0a09fff1aed940a600 Mon Sep 17 00:00:00 2001 From: rithin-pullela-aws Date: Tue, 5 May 2026 09:16:35 -0700 Subject: [PATCH 2/3] feat: Enforce name uniqueness on agent and memory-container update paths Addresses reviewer feedback on opensearch-project/ml-commons#4808: the uniqueness flags were only gating register/create paths, so a rename via PUT /agents/{id} or PUT /memory_containers/{id} could bypass the check and produce two resources with the same name even while the setting was on. Extends both update transports with a uniqueness check, gated on the same per-resource flag as the create path: - plugins.ml_commons.agent_name_uniqueness_enabled - plugins.ml_commons.agentic_memory_name_uniqueness_enabled The check runs after access validation and before the put/update. Gate shape mirrors TransportUpdateModelGroupAction#updateModelGroup: StringUtils.isBlank(newName) || newName.equals(originalName) -> skip uniqueness search so it is skipped when (a) the flag is off, (b) the update omits a new name or provides a blank one, or (c) the new name equals the current name (idempotent PUT does not 409 against itself). The ~50-line search-plumbing block that was duplicated across the four transports (TransportRegisterAgent, UpdateAgent, TransportCreate/Update MemoryContainer) is extracted into NameUniquenessHelper, mirroring the MLModelGroupManager#validateUniqueModelGroupName shape: the helper runs the tenant-scoped exact-match search and hands back the raw SearchResponse, leaving each caller to format its own 409 body and short-circuit on blank/same names. Default-off behavior is unchanged - the feature-flag gate still short-circuits before any helper call. The 409 body intentionally does NOT echo the conflicting resource's document id: doing so would let a caller confirm the existence of resources they cannot otherwise see by probing names. Only the caller-supplied name is reflected back. Regression guards (assertFalse checks) are added to the unit and integration tests for both update paths. Signed-off-by: rithin-pullela-aws --- .../agents/TransportRegisterAgentAction.java | 75 ++---- .../agents/UpdateAgentTransportAction.java | 60 ++++- .../TransportCreateMemoryContainerAction.java | 70 ++--- .../TransportUpdateMemoryContainerAction.java | 88 ++++++- .../ml/helper/NameUniquenessHelper.java | 101 +++++++ .../UpdateAgentTransportActionTests.java | 224 ++++++++++++++++ ...sportUpdateMemoryContainerActionTests.java | 247 ++++++++++++++++++ .../ml/rest/RestMLAgentNameUniquenessIT.java | 67 +++++ ...RestMLMemoryContainerNameUniquenessIT.java | 77 ++++++ 9 files changed, 898 insertions(+), 111 deletions(-) create mode 100644 plugin/src/main/java/org/opensearch/ml/helper/NameUniquenessHelper.java diff --git a/plugin/src/main/java/org/opensearch/ml/action/agents/TransportRegisterAgentAction.java b/plugin/src/main/java/org/opensearch/ml/action/agents/TransportRegisterAgentAction.java index e1e6282156..b064603d67 100644 --- a/plugin/src/main/java/org/opensearch/ml/action/agents/TransportRegisterAgentAction.java +++ b/plugin/src/main/java/org/opensearch/ml/action/agents/TransportRegisterAgentAction.java @@ -15,12 +15,10 @@ import java.util.HashMap; import java.util.Map; -import org.opensearch.ExceptionsHelper; import org.opensearch.OpenSearchException; import org.opensearch.OpenSearchStatusException; import org.opensearch.action.ActionRequest; import org.opensearch.action.index.IndexResponse; -import org.opensearch.action.search.SearchRequest; import org.opensearch.action.support.ActionFilters; import org.opensearch.action.support.HandledTransportAction; import org.opensearch.cluster.service.ClusterService; @@ -29,9 +27,6 @@ import org.opensearch.commons.authuser.User; import org.opensearch.core.action.ActionListener; import org.opensearch.core.rest.RestStatus; -import org.opensearch.index.IndexNotFoundException; -import org.opensearch.index.query.BoolQueryBuilder; -import org.opensearch.index.query.TermQueryBuilder; import org.opensearch.ml.action.agent.MLAgentRegistrationValidator; import org.opensearch.ml.action.contextmanagement.ContextManagementTemplateService; import org.opensearch.ml.common.MLAgentType; @@ -50,13 +45,12 @@ import org.opensearch.ml.engine.algorithms.agent.MLPlanExecuteAndReflectAgentRunner; import org.opensearch.ml.engine.function_calling.FunctionCallingFactory; import org.opensearch.ml.engine.indices.MLIndicesHandler; +import org.opensearch.ml.helper.NameUniquenessHelper; import org.opensearch.ml.utils.RestActionUtils; import org.opensearch.ml.utils.TenantAwareHelper; import org.opensearch.remote.metadata.client.PutDataObjectRequest; import org.opensearch.remote.metadata.client.SdkClient; -import org.opensearch.remote.metadata.client.SearchDataObjectRequest; import org.opensearch.remote.metadata.common.SdkClientUtils; -import org.opensearch.search.builder.SearchSourceBuilder; import org.opensearch.tasks.Task; import org.opensearch.transport.TransportService; import org.opensearch.transport.client.Client; @@ -159,54 +153,25 @@ private void validateAgentNameUniqueness(MLAgent mlAgent, ActionListener l } private void checkAgentNameAvailable(String name, String tenantId, ActionListener listener) { - BoolQueryBuilder query = new BoolQueryBuilder().filter(new TermQueryBuilder("name.keyword", name)); - SearchSourceBuilder sourceBuilder = new SearchSourceBuilder().query(query).size(1).fetchSource(false); - SearchRequest searchRequest = new SearchRequest(ML_AGENT_INDEX).source(sourceBuilder); - SearchDataObjectRequest searchDataObjectRequest = SearchDataObjectRequest - .builder() - .indices(searchRequest.indices()) - .searchSourceBuilder(searchRequest.source()) - .tenantId(tenantId) - .build(); - - try (ThreadContext.StoredContext context = client.threadPool().getThreadContext().stashContext()) { - sdkClient.searchDataObjectAsync(searchDataObjectRequest).whenComplete((r, throwable) -> { - context.restore(); - if (throwable != null) { - if (ExceptionsHelper.unwrap(throwable, IndexNotFoundException.class) != null) { - // Index not yet created - no duplicate possible - listener.onResponse(null); - return; - } - Exception cause = SdkClientUtils.unwrapAndConvertToException(throwable); - log.error("Failed to search ML agent index for name uniqueness check", cause); - listener.onFailure(cause); - return; - } - try { - long totalHits = r.searchResponse().getHits().getTotalHits() == null - ? 0 - : r.searchResponse().getHits().getTotalHits().value(); - if (totalHits > 0) { - listener - .onFailure( - new OpenSearchStatusException( - "An agent with name [" + name + "] already exists. Agent names must be unique.", - RestStatus.CONFLICT - ) - ); - } else { - listener.onResponse(null); - } - } catch (Exception e) { - log.error("Failed to parse search response for agent name uniqueness check", e); - listener.onFailure(e); - } - }); - } catch (Exception e) { - log.error("Failed to execute agent name uniqueness check", e); - listener.onFailure(e); - } + NameUniquenessHelper.searchByExactName(client, sdkClient, ML_AGENT_INDEX, name, tenantId, ActionListener.wrap(response -> { + if (response == null) { + // Index not yet created - no duplicate possible. + listener.onResponse(null); + return; + } + long totalHits = response.getHits().getTotalHits() == null ? 0 : response.getHits().getTotalHits().value(); + if (totalHits > 0) { + listener + .onFailure( + new OpenSearchStatusException( + "An agent with name [" + name + "] already exists. Agent names must be unique.", + RestStatus.CONFLICT + ) + ); + } else { + listener.onResponse(null); + } + }, listener::onFailure)); } private void createModelAndRegisterAgent(MLAgent mlAgent, ActionListener listener) { diff --git a/plugin/src/main/java/org/opensearch/ml/action/agents/UpdateAgentTransportAction.java b/plugin/src/main/java/org/opensearch/ml/action/agents/UpdateAgentTransportAction.java index 363a91d624..758f9f1538 100644 --- a/plugin/src/main/java/org/opensearch/ml/action/agents/UpdateAgentTransportAction.java +++ b/plugin/src/main/java/org/opensearch/ml/action/agents/UpdateAgentTransportAction.java @@ -9,6 +9,7 @@ import java.time.Instant; +import org.apache.commons.lang3.StringUtils; import org.opensearch.OpenSearchStatusException; import org.opensearch.action.ActionRequest; import org.opensearch.action.DocWriteResponse; @@ -32,6 +33,7 @@ import org.opensearch.ml.common.transport.agent.MLAgentUpdateAction; import org.opensearch.ml.common.transport.agent.MLAgentUpdateInput; import org.opensearch.ml.common.transport.agent.MLAgentUpdateRequest; +import org.opensearch.ml.helper.NameUniquenessHelper; import org.opensearch.ml.utils.RestActionUtils; import org.opensearch.ml.utils.TenantAwareHelper; import org.opensearch.remote.metadata.client.GetDataObjectRequest; @@ -127,7 +129,15 @@ protected void doExecute(Task task, ActionRequest request, ActionListener updateAgent(agentId, mlAgentUpdateInput, retrievedAgent, wrappedListener), + wrappedListener::onFailure + ) + ); } } else { log.error("Failed to validate tenant for Agent ID {}", agentId); @@ -180,6 +190,54 @@ private void updateAgent( }); } + /** + * When {@code plugins.ml_commons.agent_name_uniqueness_enabled} is enabled and the update + * would rename the agent to a new value, reject the update if another agent in the same + * tenant already has that name. Skipped if the setting is off, if the update omits a new + * name (or provides a blank one), or if the new name equals the existing name (so a PUT + * that doesn't actually rename does not 409 against itself). Mirrors the gate shape used + * by {@code TransportUpdateModelGroupAction#updateModelGroup}. + * + *

Same best-effort semantics as the register path: two concurrent updates racing on the + * same name can both see zero hits and both succeed. + */ + private void validateAgentNameUniquenessForUpdate( + MLAgentUpdateInput updateInput, + MLAgent originalAgent, + ActionListener listener + ) { + if (!mlFeatureEnabledSetting.isAgentNameUniquenessEnabled()) { + listener.onResponse(null); + return; + } + String newName = updateInput.getName(); + if (StringUtils.isBlank(newName) || newName.equals(originalAgent.getName())) { + // No rename (name omitted/blank), or a no-op rename to the current name. + listener.onResponse(null); + return; + } + + NameUniquenessHelper + .searchByExactName(client, sdkClient, ML_AGENT_INDEX, newName, updateInput.getTenantId(), ActionListener.wrap(response -> { + if (response == null) { + listener.onResponse(null); + return; + } + long totalHits = response.getHits().getTotalHits() == null ? 0 : response.getHits().getTotalHits().value(); + if (totalHits > 0) { + listener + .onFailure( + new OpenSearchStatusException( + "An agent with name [" + newName + "] already exists. Agent names must be unique.", + RestStatus.CONFLICT + ) + ); + } else { + listener.onResponse(null); + } + }, listener::onFailure)); + } + @VisibleForTesting boolean isSuperAdminUserWrapper(ClusterService clusterService, Client client) { return RestActionUtils.isSuperAdminUser(clusterService, client); diff --git a/plugin/src/main/java/org/opensearch/ml/action/memorycontainer/TransportCreateMemoryContainerAction.java b/plugin/src/main/java/org/opensearch/ml/action/memorycontainer/TransportCreateMemoryContainerAction.java index 70212643c0..d637fcecc6 100644 --- a/plugin/src/main/java/org/opensearch/ml/action/memorycontainer/TransportCreateMemoryContainerAction.java +++ b/plugin/src/main/java/org/opensearch/ml/action/memorycontainer/TransportCreateMemoryContainerAction.java @@ -10,11 +10,9 @@ import java.time.Instant; -import org.opensearch.ExceptionsHelper; import org.opensearch.OpenSearchStatusException; import org.opensearch.action.DocWriteResponse; import org.opensearch.action.index.IndexResponse; -import org.opensearch.action.search.SearchRequest; import org.opensearch.action.support.ActionFilters; import org.opensearch.action.support.HandledTransportAction; import org.opensearch.action.support.WriteRequest; @@ -23,9 +21,6 @@ import org.opensearch.commons.authuser.User; import org.opensearch.core.action.ActionListener; import org.opensearch.core.rest.RestStatus; -import org.opensearch.index.IndexNotFoundException; -import org.opensearch.index.query.BoolQueryBuilder; -import org.opensearch.index.query.TermQueryBuilder; import org.opensearch.ml.common.memorycontainer.MLMemoryContainer; import org.opensearch.ml.common.memorycontainer.MemoryConfiguration; import org.opensearch.ml.common.memorycontainer.MemoryStrategy; @@ -39,14 +34,13 @@ import org.opensearch.ml.helper.MemoryContainerModelValidator; import org.opensearch.ml.helper.MemoryContainerPipelineHelper; import org.opensearch.ml.helper.MemoryContainerSharedIndexValidator; +import org.opensearch.ml.helper.NameUniquenessHelper; import org.opensearch.ml.model.MLModelManager; import org.opensearch.ml.utils.RestActionUtils; import org.opensearch.ml.utils.TenantAwareHelper; import org.opensearch.remote.metadata.client.PutDataObjectRequest; import org.opensearch.remote.metadata.client.SdkClient; -import org.opensearch.remote.metadata.client.SearchDataObjectRequest; import org.opensearch.remote.metadata.common.SdkClientUtils; -import org.opensearch.search.builder.SearchSourceBuilder; import org.opensearch.tasks.Task; import org.opensearch.transport.TransportService; import org.opensearch.transport.client.Client; @@ -319,54 +313,26 @@ private void validateMemoryContainerNameUniqueness(String name, String tenantId, // MLCreateMemoryContainerInput rejects null names upstream, so we rely on that invariant // here and don't re-check. - BoolQueryBuilder query = new BoolQueryBuilder().filter(new TermQueryBuilder("name.keyword", name)); - SearchSourceBuilder sourceBuilder = new SearchSourceBuilder().query(query).size(1).fetchSource(false); - SearchRequest searchRequest = new SearchRequest(ML_MEMORY_CONTAINER_INDEX).source(sourceBuilder); - SearchDataObjectRequest searchDataObjectRequest = SearchDataObjectRequest - .builder() - .indices(searchRequest.indices()) - .searchSourceBuilder(searchRequest.source()) - .tenantId(tenantId) - .build(); - - try (ThreadContext.StoredContext context = client.threadPool().getThreadContext().stashContext()) { - sdkClient.searchDataObjectAsync(searchDataObjectRequest).whenComplete((r, throwable) -> { - context.restore(); - if (throwable != null) { - if (ExceptionsHelper.unwrap(throwable, IndexNotFoundException.class) != null) { - // Index not yet created - no duplicate possible - listener.onResponse(null); - return; - } - Exception cause = SdkClientUtils.unwrapAndConvertToException(throwable); - log.error("Failed to search memory container index for name uniqueness check", cause); - listener.onFailure(cause); + NameUniquenessHelper + .searchByExactName(client, sdkClient, ML_MEMORY_CONTAINER_INDEX, name, tenantId, ActionListener.wrap(response -> { + if (response == null) { + // Index not yet created - no duplicate possible. + listener.onResponse(null); return; } - try { - long totalHits = r.searchResponse().getHits().getTotalHits() == null - ? 0 - : r.searchResponse().getHits().getTotalHits().value(); - if (totalHits > 0) { - listener - .onFailure( - new OpenSearchStatusException( - "A memory container with name [" + name + "] already exists. Memory container names must be unique.", - RestStatus.CONFLICT - ) - ); - } else { - listener.onResponse(null); - } - } catch (Exception e) { - log.error("Failed to parse search response for memory container name uniqueness check", e); - listener.onFailure(e); + long totalHits = response.getHits().getTotalHits() == null ? 0 : response.getHits().getTotalHits().value(); + if (totalHits > 0) { + listener + .onFailure( + new OpenSearchStatusException( + "A memory container with name [" + name + "] already exists. Memory container names must be unique.", + RestStatus.CONFLICT + ) + ); + } else { + listener.onResponse(null); } - }); - } catch (Exception e) { - log.error("Failed to execute memory container name uniqueness check", e); - listener.onFailure(e); - } + }, listener::onFailure)); } private void validateConfiguration(MemoryConfiguration config, ActionListener listener) { diff --git a/plugin/src/main/java/org/opensearch/ml/action/memorycontainer/TransportUpdateMemoryContainerAction.java b/plugin/src/main/java/org/opensearch/ml/action/memorycontainer/TransportUpdateMemoryContainerAction.java index 72ad3e32de..52efd06f5c 100644 --- a/plugin/src/main/java/org/opensearch/ml/action/memorycontainer/TransportUpdateMemoryContainerAction.java +++ b/plugin/src/main/java/org/opensearch/ml/action/memorycontainer/TransportUpdateMemoryContainerAction.java @@ -17,6 +17,7 @@ import java.util.List; import java.util.Map; +import org.apache.commons.lang3.StringUtils; import org.opensearch.OpenSearchStatusException; import org.opensearch.action.ActionRequest; import org.opensearch.action.support.ActionFilters; @@ -42,6 +43,7 @@ import org.opensearch.ml.helper.MemoryContainerModelValidator; import org.opensearch.ml.helper.MemoryContainerPipelineHelper; import org.opensearch.ml.helper.MemoryContainerSharedIndexValidator; +import org.opensearch.ml.helper.NameUniquenessHelper; import org.opensearch.ml.helper.StrategyMergeHelper; import org.opensearch.ml.model.MLModelManager; import org.opensearch.ml.utils.RestActionUtils; @@ -123,6 +125,33 @@ protected void doExecute(Task task, ActionRequest request, ActionListener { + continueUpdate( + container, + newName, + newDescription, + allowedBackendRoles, + updateConfiguration, + memoryContainerId, + actionListener + ); + }, actionListener::onFailure)); + }, e -> { + log.error("Failed to get memory container for update", e); + actionListener.onFailure(new OpenSearchStatusException("Internal server error", RestStatus.INTERNAL_SERVER_ERROR)); + })); + } + + private void continueUpdate( + MLMemoryContainer container, + String newName, + String newDescription, + List allowedBackendRoles, + MemoryConfiguration updateConfiguration, + String memoryContainerId, + ActionListener actionListener + ) { + try { // Prepare the update Map updateFields = new HashMap<>(); if (newName != null) { @@ -195,10 +224,63 @@ protected void doExecute(Task task, ActionRequest request, ActionListener { - log.error("Failed to get memory container for update", e); + } catch (Exception e) { + log.error("Failed to update memory container {}", memoryContainerId, e); actionListener.onFailure(new OpenSearchStatusException("Internal server error", RestStatus.INTERNAL_SERVER_ERROR)); - })); + } + } + + /** + * When {@code plugins.ml_commons.agentic_memory_name_uniqueness_enabled} is enabled and the + * update would rename the container to a new value, reject the update if another memory + * container in the same tenant already has that name. Skipped if the setting is off, if the + * update omits a new name (or provides a blank one), or if the new name equals the existing + * name. Mirrors the gate shape used by {@code TransportUpdateModelGroupAction#updateModelGroup}. + * + *

Same best-effort semantics as the create path: two concurrent updates racing on the + * same name can both see zero hits and both succeed. + */ + private void validateMemoryContainerNameUniquenessForUpdate( + String newName, + MLMemoryContainer originalContainer, + ActionListener listener + ) { + if (!mlFeatureEnabledSetting.isAgenticMemoryNameUniquenessEnabled()) { + listener.onResponse(null); + return; + } + if (StringUtils.isBlank(newName) || newName.equals(originalContainer.getName())) { + // No rename (name omitted/blank), or a no-op rename to the current name. + listener.onResponse(null); + return; + } + + NameUniquenessHelper + .searchByExactName( + client, + sdkClient, + ML_MEMORY_CONTAINER_INDEX, + newName, + originalContainer.getTenantId(), + ActionListener.wrap(response -> { + if (response == null) { + listener.onResponse(null); + return; + } + long totalHits = response.getHits().getTotalHits() == null ? 0 : response.getHits().getTotalHits().value(); + if (totalHits > 0) { + listener + .onFailure( + new OpenSearchStatusException( + "A memory container with name [" + newName + "] already exists. Memory container names must be unique.", + RestStatus.CONFLICT + ) + ); + } else { + listener.onResponse(null); + } + }, listener::onFailure) + ); } /** diff --git a/plugin/src/main/java/org/opensearch/ml/helper/NameUniquenessHelper.java b/plugin/src/main/java/org/opensearch/ml/helper/NameUniquenessHelper.java new file mode 100644 index 0000000000..63d7a6bc5e --- /dev/null +++ b/plugin/src/main/java/org/opensearch/ml/helper/NameUniquenessHelper.java @@ -0,0 +1,101 @@ +/* + * Copyright OpenSearch Contributors + * SPDX-License-Identifier: Apache-2.0 + */ + +package org.opensearch.ml.helper; + +import org.opensearch.ExceptionsHelper; +import org.opensearch.action.search.SearchRequest; +import org.opensearch.action.search.SearchResponse; +import org.opensearch.common.util.concurrent.ThreadContext; +import org.opensearch.core.action.ActionListener; +import org.opensearch.index.IndexNotFoundException; +import org.opensearch.index.query.BoolQueryBuilder; +import org.opensearch.index.query.TermQueryBuilder; +import org.opensearch.remote.metadata.client.SdkClient; +import org.opensearch.remote.metadata.client.SearchDataObjectRequest; +import org.opensearch.remote.metadata.common.SdkClientUtils; +import org.opensearch.search.builder.SearchSourceBuilder; +import org.opensearch.transport.client.Client; + +import lombok.extern.log4j.Log4j2; + +/** + * Shared plumbing for "is this name already used in a tenant-scoped index?" checks used by the + * agent and memory-container uniqueness flags. Mirrors the pattern in + * {@code MLModelGroupManager#validateUniqueModelGroupName}: runs an exact-match query on the + * {@code name.keyword} subfield under the tenant's scope and hands the raw {@link SearchResponse} + * back to the caller, which is then responsible for interpreting hits, formatting the 409 + * message, and short-circuiting on blank/same names. + * + *

If the target index has not been created yet, the listener receives {@code null} (treated + * as "no conflict possible") rather than an error. + * + *

Do not echo the conflicting hit's {@code _id} in the 409 response body. Doing so + * would let a caller confirm the existence of resources they cannot otherwise see by probing + * names. The name itself is always caller-supplied input, so echoing it back is safe. + */ +@Log4j2 +public final class NameUniquenessHelper { + + private NameUniquenessHelper() {} + + /** + * Search an index for documents whose {@code name.keyword} exactly matches {@code name}, scoped + * to the given tenant. + * + * @param client node client (used only for its ThreadContext - the search itself goes + * through {@code sdkClient}) + * @param sdkClient SDK client used to issue the tenant-scoped search + * @param indexName index to query + * @param name value to match on {@code name.keyword} + * @param tenantId tenant scope for the search (may be null in single-tenant mode) + * @param listener receives the {@link SearchResponse} on success, {@code null} if the index + * does not exist, or a failure otherwise + */ + public static void searchByExactName( + Client client, + SdkClient sdkClient, + String indexName, + String name, + String tenantId, + ActionListener listener + ) { + BoolQueryBuilder query = new BoolQueryBuilder().filter(new TermQueryBuilder("name.keyword", name)); + SearchSourceBuilder sourceBuilder = new SearchSourceBuilder().query(query).size(1).fetchSource(false); + SearchRequest searchRequest = new SearchRequest(indexName).source(sourceBuilder); + SearchDataObjectRequest searchDataObjectRequest = SearchDataObjectRequest + .builder() + .indices(searchRequest.indices()) + .searchSourceBuilder(searchRequest.source()) + .tenantId(tenantId) + .build(); + + try (ThreadContext.StoredContext context = client.threadPool().getThreadContext().stashContext()) { + sdkClient.searchDataObjectAsync(searchDataObjectRequest).whenComplete((r, throwable) -> { + context.restore(); + if (throwable != null) { + if (ExceptionsHelper.unwrap(throwable, IndexNotFoundException.class) != null) { + // Index not yet created - no duplicate possible. + listener.onResponse(null); + return; + } + Exception cause = SdkClientUtils.unwrapAndConvertToException(throwable); + log.error("Failed to search index [{}] for name uniqueness check", indexName, cause); + listener.onFailure(cause); + return; + } + try { + listener.onResponse(r.searchResponse()); + } catch (Exception e) { + log.error("Failed to parse search response for name uniqueness check on index [{}]", indexName, e); + listener.onFailure(e); + } + }); + } catch (Exception e) { + log.error("Failed to execute name uniqueness check on index [{}]", indexName, e); + listener.onFailure(e); + } + } +} diff --git a/plugin/src/test/java/org/opensearch/ml/action/agents/UpdateAgentTransportActionTests.java b/plugin/src/test/java/org/opensearch/ml/action/agents/UpdateAgentTransportActionTests.java index fe665ed16d..1fa6f1e343 100644 --- a/plugin/src/test/java/org/opensearch/ml/action/agents/UpdateAgentTransportActionTests.java +++ b/plugin/src/test/java/org/opensearch/ml/action/agents/UpdateAgentTransportActionTests.java @@ -6,6 +6,8 @@ package org.opensearch.ml.action.agents; import static org.junit.Assert.assertEquals; +import static org.junit.Assert.assertFalse; +import static org.junit.Assert.assertTrue; import static org.mockito.ArgumentMatchers.any; import static org.mockito.Mockito.*; import static org.opensearch.ml.common.CommonValue.ML_AGENT_INDEX; @@ -290,6 +292,228 @@ public void testDoExecute_HiddenAgentNonSuperAdmin() throws IOException { assertEquals(RestStatus.FORBIDDEN, argumentCaptor.getValue().status()); } + @Test + public void testDoExecute_uniquenessEnforced_renameToExistingName_rejected() throws IOException { + when(mlFeatureEnabledSetting.isAgentNameUniquenessEnabled()).thenReturn(true); + + String agentId = "test_agent_id"; + MLAgentUpdateInput mlAgentUpdateInput = MLAgentUpdateInput + .builder() + .agentId(agentId) + .name("other_agent") + .description("desc") + .build(); + + GetResponse getResponse = prepareMLAgentGetResponse(agentId, false, null); + + MLAgentUpdateRequest updateRequest = mock(MLAgentUpdateRequest.class); + when(updateRequest.getMlAgentUpdateInput()).thenReturn(mlAgentUpdateInput); + doReturn(true).when(updateAgentTransportAction).isSuperAdminUserWrapper(clusterService, client); + + doAnswer(invocation -> { + ActionListener listener = invocation.getArgument(1); + listener.onResponse(getResponse); + return null; + }).when(client).get(any(), any()); + + doAnswer(invocation -> { + ActionListener al = invocation.getArgument(1); + al.onResponse(buildSearchResponseWithConflictingHit("conflicting_agent_id")); + return null; + }).when(client).search(any(), any()); + + updateAgentTransportAction.doExecute(task, updateRequest, actionListener); + + ArgumentCaptor captor = ArgumentCaptor.forClass(Exception.class); + verify(actionListener).onFailure(captor.capture()); + assertTrue(captor.getValue() instanceof OpenSearchStatusException); + assertEquals(RestStatus.CONFLICT, ((OpenSearchStatusException) captor.getValue()).status()); + assertTrue(captor.getValue().getMessage().contains("already exists")); + assertTrue(captor.getValue().getMessage().contains("other_agent")); + assertFalse(captor.getValue().getMessage().contains("conflicting_agent_id")); + // The put must not have been issued + verify(client, times(0)).update(any(), any()); + } + + @Test + public void testDoExecute_uniquenessEnforced_renameToUnusedName_allowed() throws IOException { + when(mlFeatureEnabledSetting.isAgentNameUniquenessEnabled()).thenReturn(true); + + String agentId = "test_agent_id"; + MLAgentUpdateInput mlAgentUpdateInput = MLAgentUpdateInput + .builder() + .agentId(agentId) + .name("brand_new_name") + .description("desc") + .build(); + + GetResponse getResponse = prepareMLAgentGetResponse(agentId, false, null); + + MLAgentUpdateRequest updateRequest = mock(MLAgentUpdateRequest.class); + when(updateRequest.getMlAgentUpdateInput()).thenReturn(mlAgentUpdateInput); + doReturn(true).when(updateAgentTransportAction).isSuperAdminUserWrapper(clusterService, client); + + doAnswer(invocation -> { + ActionListener listener = invocation.getArgument(1); + listener.onResponse(getResponse); + return null; + }).when(client).get(any(), any()); + + doAnswer(invocation -> { + ActionListener al = invocation.getArgument(1); + al.onResponse(buildSearchResponseWithHits(0L)); + return null; + }).when(client).search(any(), any()); + + doAnswer(invocation -> { + ActionListener listener = invocation.getArgument(1); + listener.onResponse(updateResponse); + return null; + }).when(client).update(any(), any()); + + updateAgentTransportAction.doExecute(task, updateRequest, actionListener); + + ArgumentCaptor captor = ArgumentCaptor.forClass(UpdateResponse.class); + verify(actionListener).onResponse(captor.capture()); + assertEquals(DocWriteResponse.Result.UPDATED, captor.getValue().getResult()); + } + + @Test + public void testDoExecute_uniquenessEnforced_sameNameNoOp_skipsSearch() throws IOException { + when(mlFeatureEnabledSetting.isAgentNameUniquenessEnabled()).thenReturn(true); + + String agentId = "test_agent_id"; + // prepareMLAgentGetResponse stores name="test"; PUT with the same name must be a no-op. + MLAgentUpdateInput mlAgentUpdateInput = MLAgentUpdateInput + .builder() + .agentId(agentId) + .name("test") + .description("desc") + .build(); + + GetResponse getResponse = prepareMLAgentGetResponse(agentId, false, null); + + MLAgentUpdateRequest updateRequest = mock(MLAgentUpdateRequest.class); + when(updateRequest.getMlAgentUpdateInput()).thenReturn(mlAgentUpdateInput); + doReturn(true).when(updateAgentTransportAction).isSuperAdminUserWrapper(clusterService, client); + + doAnswer(invocation -> { + ActionListener listener = invocation.getArgument(1); + listener.onResponse(getResponse); + return null; + }).when(client).get(any(), any()); + + doAnswer(invocation -> { + ActionListener listener = invocation.getArgument(1); + listener.onResponse(updateResponse); + return null; + }).when(client).update(any(), any()); + + updateAgentTransportAction.doExecute(task, updateRequest, actionListener); + + verify(client, times(0)).search(any(), any()); + verify(actionListener).onResponse(any()); + } + + @Test + public void testDoExecute_uniquenessDisabled_renameSkipsSearch() throws IOException { + when(mlFeatureEnabledSetting.isAgentNameUniquenessEnabled()).thenReturn(false); + + String agentId = "test_agent_id"; + MLAgentUpdateInput mlAgentUpdateInput = MLAgentUpdateInput + .builder() + .agentId(agentId) + .name("renamed_but_flag_off") + .description("desc") + .build(); + + GetResponse getResponse = prepareMLAgentGetResponse(agentId, false, null); + + MLAgentUpdateRequest updateRequest = mock(MLAgentUpdateRequest.class); + when(updateRequest.getMlAgentUpdateInput()).thenReturn(mlAgentUpdateInput); + doReturn(true).when(updateAgentTransportAction).isSuperAdminUserWrapper(clusterService, client); + + doAnswer(invocation -> { + ActionListener listener = invocation.getArgument(1); + listener.onResponse(getResponse); + return null; + }).when(client).get(any(), any()); + + doAnswer(invocation -> { + ActionListener listener = invocation.getArgument(1); + listener.onResponse(updateResponse); + return null; + }).when(client).update(any(), any()); + + updateAgentTransportAction.doExecute(task, updateRequest, actionListener); + + verify(client, times(0)).search(any(), any()); + verify(actionListener).onResponse(any()); + } + + private org.opensearch.action.search.SearchResponse buildSearchResponseWithHits(long hitCount) { + org.apache.lucene.search.TotalHits totalHits = new org.apache.lucene.search.TotalHits( + hitCount, + org.apache.lucene.search.TotalHits.Relation.EQUAL_TO + ); + org.opensearch.search.SearchHits hits = new org.opensearch.search.SearchHits( + new org.opensearch.search.SearchHit[0], + totalHits, + Float.NaN + ); + org.opensearch.search.internal.InternalSearchResponse internal = new org.opensearch.search.internal.InternalSearchResponse( + hits, + org.opensearch.search.aggregations.InternalAggregations.EMPTY, + null, + null, + false, + null, + 0 + ); + return new org.opensearch.action.search.SearchResponse( + internal, + null, + 1, + 1, + 0, + 1, + org.opensearch.action.search.ShardSearchFailure.EMPTY_ARRAY, + org.opensearch.action.search.SearchResponse.Clusters.EMPTY + ); + } + + private org.opensearch.action.search.SearchResponse buildSearchResponseWithConflictingHit(String conflictingId) { + org.apache.lucene.search.TotalHits totalHits = new org.apache.lucene.search.TotalHits( + 1L, + org.apache.lucene.search.TotalHits.Relation.EQUAL_TO + ); + org.opensearch.search.SearchHit hit = new org.opensearch.search.SearchHit(0, conflictingId, null, null); + org.opensearch.search.SearchHits hits = new org.opensearch.search.SearchHits( + new org.opensearch.search.SearchHit[] { hit }, + totalHits, + Float.NaN + ); + org.opensearch.search.internal.InternalSearchResponse internal = new org.opensearch.search.internal.InternalSearchResponse( + hits, + org.opensearch.search.aggregations.InternalAggregations.EMPTY, + null, + null, + false, + null, + 0 + ); + return new org.opensearch.action.search.SearchResponse( + internal, + null, + 1, + 1, + 0, + 1, + org.opensearch.action.search.ShardSearchFailure.EMPTY_ARRAY, + org.opensearch.action.search.SearchResponse.Clusters.EMPTY + ); + } + private GetResponse prepareMLAgentGetResponse(String agentId, boolean isHidden, String tenantId) throws IOException { MLAgent mlAgent = MLAgent .builder() diff --git a/plugin/src/test/java/org/opensearch/ml/action/memorycontainer/TransportUpdateMemoryContainerActionTests.java b/plugin/src/test/java/org/opensearch/ml/action/memorycontainer/TransportUpdateMemoryContainerActionTests.java index 6cdaa57f4e..d18690da6d 100644 --- a/plugin/src/test/java/org/opensearch/ml/action/memorycontainer/TransportUpdateMemoryContainerActionTests.java +++ b/plugin/src/test/java/org/opensearch/ml/action/memorycontainer/TransportUpdateMemoryContainerActionTests.java @@ -2283,4 +2283,251 @@ public void testUpdateContainer_TransitionToStrategies_LongTermIndexCreationFail verify(listener, never()).onResponse(any()); } + public void testDoExecute_uniquenessEnforced_renameToExistingName_rejected() { + when(mlFeatureEnabledSetting.isAgenticMemoryNameUniquenessEnabled()).thenReturn(true); + + String containerId = "test-container-id"; + String newName = "taken-name"; + + MLUpdateMemoryContainerInput input = MLUpdateMemoryContainerInput.builder().name(newName).build(); + MLUpdateMemoryContainerRequest request = MLUpdateMemoryContainerRequest + .builder() + .memoryContainerId(containerId) + .mlUpdateMemoryContainerInput(input) + .build(); + + ActionListener listener = mock(ActionListener.class); + + MLMemoryContainer container = MLMemoryContainer + .builder() + .name("old-name") + .owner(new User("test-user", Collections.emptyList(), Collections.emptyList(), Collections.emptyMap())) + .build(); + + doAnswer(invocation -> { + ActionListener containerListener = invocation.getArgument(1); + containerListener.onResponse(container); + return null; + }).when(memoryContainerHelper).getMemoryContainer(any(), any()); + when(memoryContainerHelper.checkMemoryContainerAccess(isNull(), eq(container))).thenReturn(true); + + java.util.concurrent.CompletableFuture< + org.opensearch.remote.metadata.client.SearchDataObjectResponse + > future = java.util.concurrent.CompletableFuture.completedFuture(searchDataObjectResponseWithHit("conflicting-container-id")); + when(sdkClient.searchDataObjectAsync(any(org.opensearch.remote.metadata.client.SearchDataObjectRequest.class))).thenReturn(future); + + action.doExecute(task, request, listener); + + ArgumentCaptor captor = ArgumentCaptor.forClass(Exception.class); + verify(listener).onFailure(captor.capture()); + Exception exception = captor.getValue(); + assertTrue(exception instanceof OpenSearchStatusException); + assertEquals(RestStatus.CONFLICT, ((OpenSearchStatusException) exception).status()); + assertTrue(exception.getMessage().contains("already exists")); + assertTrue(exception.getMessage().contains(newName)); + assertFalse(exception.getMessage().contains("conflicting-container-id")); + // Must not have attempted a write when a duplicate is found. + verify(client, never()).update(any(), any()); + } + + public void testDoExecute_uniquenessEnforced_renameToUnusedName_allowed() { + when(mlFeatureEnabledSetting.isAgenticMemoryNameUniquenessEnabled()).thenReturn(true); + + String containerId = "test-container-id"; + String newName = "brand-new-name"; + + MLUpdateMemoryContainerInput input = MLUpdateMemoryContainerInput.builder().name(newName).build(); + MLUpdateMemoryContainerRequest request = MLUpdateMemoryContainerRequest + .builder() + .memoryContainerId(containerId) + .mlUpdateMemoryContainerInput(input) + .build(); + + ActionListener listener = mock(ActionListener.class); + + MLMemoryContainer container = MLMemoryContainer + .builder() + .name("old-name") + .owner(new User("test-user", Collections.emptyList(), Collections.emptyList(), Collections.emptyMap())) + .build(); + + doAnswer(invocation -> { + ActionListener containerListener = invocation.getArgument(1); + containerListener.onResponse(container); + return null; + }).when(memoryContainerHelper).getMemoryContainer(any(), any()); + when(memoryContainerHelper.checkMemoryContainerAccess(isNull(), eq(container))).thenReturn(true); + + java.util.concurrent.CompletableFuture< + org.opensearch.remote.metadata.client.SearchDataObjectResponse + > future = java.util.concurrent.CompletableFuture.completedFuture(searchDataObjectResponseWithHits(0L)); + when(sdkClient.searchDataObjectAsync(any(org.opensearch.remote.metadata.client.SearchDataObjectRequest.class))).thenReturn(future); + + UpdateResponse updateResponse = new UpdateResponse( + new ShardId(new Index("test", "uuid"), 0), + containerId, + 1L, + 1L, + 1L, + org.opensearch.action.DocWriteResponse.Result.UPDATED + ); + doAnswer(invocation -> { + ActionListener updateListener = invocation.getArgument(1); + updateListener.onResponse(updateResponse); + return null; + }).when(client).update(any(), any()); + + action.doExecute(task, request, listener); + + verify(sdkClient).searchDataObjectAsync(any(org.opensearch.remote.metadata.client.SearchDataObjectRequest.class)); + verify(listener).onResponse(updateResponse); + } + + public void testDoExecute_uniquenessEnforced_sameNameNoOp_skipsSearch() { + when(mlFeatureEnabledSetting.isAgenticMemoryNameUniquenessEnabled()).thenReturn(true); + + String containerId = "test-container-id"; + // A PUT that re-specifies the existing name must be a no-op rename (no search, no 409). + MLUpdateMemoryContainerInput input = MLUpdateMemoryContainerInput.builder().name("old-name").description("new desc").build(); + MLUpdateMemoryContainerRequest request = MLUpdateMemoryContainerRequest + .builder() + .memoryContainerId(containerId) + .mlUpdateMemoryContainerInput(input) + .build(); + + ActionListener listener = mock(ActionListener.class); + + MLMemoryContainer container = MLMemoryContainer + .builder() + .name("old-name") + .owner(new User("test-user", Collections.emptyList(), Collections.emptyList(), Collections.emptyMap())) + .build(); + + doAnswer(invocation -> { + ActionListener containerListener = invocation.getArgument(1); + containerListener.onResponse(container); + return null; + }).when(memoryContainerHelper).getMemoryContainer(any(), any()); + when(memoryContainerHelper.checkMemoryContainerAccess(isNull(), eq(container))).thenReturn(true); + + UpdateResponse updateResponse = new UpdateResponse( + new ShardId(new Index("test", "uuid"), 0), + containerId, + 1L, + 1L, + 1L, + org.opensearch.action.DocWriteResponse.Result.UPDATED + ); + doAnswer(invocation -> { + ActionListener updateListener = invocation.getArgument(1); + updateListener.onResponse(updateResponse); + return null; + }).when(client).update(any(), any()); + + action.doExecute(task, request, listener); + + verify(sdkClient, never()).searchDataObjectAsync(any(org.opensearch.remote.metadata.client.SearchDataObjectRequest.class)); + verify(listener).onResponse(updateResponse); + } + + public void testDoExecute_uniquenessDisabled_renameSkipsSearch() { + when(mlFeatureEnabledSetting.isAgenticMemoryNameUniquenessEnabled()).thenReturn(false); + + String containerId = "test-container-id"; + MLUpdateMemoryContainerInput input = MLUpdateMemoryContainerInput.builder().name("renamed-but-flag-off").build(); + MLUpdateMemoryContainerRequest request = MLUpdateMemoryContainerRequest + .builder() + .memoryContainerId(containerId) + .mlUpdateMemoryContainerInput(input) + .build(); + + ActionListener listener = mock(ActionListener.class); + + MLMemoryContainer container = MLMemoryContainer + .builder() + .name("old-name") + .owner(new User("test-user", Collections.emptyList(), Collections.emptyList(), Collections.emptyMap())) + .build(); + + doAnswer(invocation -> { + ActionListener containerListener = invocation.getArgument(1); + containerListener.onResponse(container); + return null; + }).when(memoryContainerHelper).getMemoryContainer(any(), any()); + when(memoryContainerHelper.checkMemoryContainerAccess(isNull(), eq(container))).thenReturn(true); + + UpdateResponse updateResponse = new UpdateResponse( + new ShardId(new Index("test", "uuid"), 0), + containerId, + 1L, + 1L, + 1L, + org.opensearch.action.DocWriteResponse.Result.UPDATED + ); + doAnswer(invocation -> { + ActionListener updateListener = invocation.getArgument(1); + updateListener.onResponse(updateResponse); + return null; + }).when(client).update(any(), any()); + + action.doExecute(task, request, listener); + + verify(sdkClient, never()).searchDataObjectAsync(any(org.opensearch.remote.metadata.client.SearchDataObjectRequest.class)); + verify(listener).onResponse(updateResponse); + } + + private org.opensearch.remote.metadata.client.SearchDataObjectResponse searchDataObjectResponseWithHits(long hitCount) { + org.apache.lucene.search.TotalHits totalHits = new org.apache.lucene.search.TotalHits( + hitCount, + org.apache.lucene.search.TotalHits.Relation.EQUAL_TO + ); + org.opensearch.search.SearchHits hits = new org.opensearch.search.SearchHits( + new org.opensearch.search.SearchHit[0], + totalHits, + Float.NaN + ); + return buildSdkSearchResponse(hits); + } + + private org.opensearch.remote.metadata.client.SearchDataObjectResponse searchDataObjectResponseWithHit(String conflictingId) { + org.apache.lucene.search.TotalHits totalHits = new org.apache.lucene.search.TotalHits( + 1L, + org.apache.lucene.search.TotalHits.Relation.EQUAL_TO + ); + org.opensearch.search.SearchHit hit = new org.opensearch.search.SearchHit(0, conflictingId, null, null); + org.opensearch.search.SearchHits hits = new org.opensearch.search.SearchHits( + new org.opensearch.search.SearchHit[] { hit }, + totalHits, + Float.NaN + ); + return buildSdkSearchResponse(hits); + } + + private org.opensearch.remote.metadata.client.SearchDataObjectResponse buildSdkSearchResponse(org.opensearch.search.SearchHits hits) { + org.opensearch.search.internal.InternalSearchResponse internal = new org.opensearch.search.internal.InternalSearchResponse( + hits, + org.opensearch.search.aggregations.InternalAggregations.EMPTY, + null, + null, + false, + null, + 0 + ); + org.opensearch.action.search.SearchResponse searchResponse = new org.opensearch.action.search.SearchResponse( + internal, + null, + 1, + 1, + 0, + 1, + org.opensearch.action.search.ShardSearchFailure.EMPTY_ARRAY, + org.opensearch.action.search.SearchResponse.Clusters.EMPTY + ); + org.opensearch.remote.metadata.client.SearchDataObjectResponse sdkResp = mock( + org.opensearch.remote.metadata.client.SearchDataObjectResponse.class + ); + when(sdkResp.searchResponse()).thenReturn(searchResponse); + return sdkResp; + } + } diff --git a/plugin/src/test/java/org/opensearch/ml/rest/RestMLAgentNameUniquenessIT.java b/plugin/src/test/java/org/opensearch/ml/rest/RestMLAgentNameUniquenessIT.java index 64fab72d75..af28a806fa 100644 --- a/plugin/src/test/java/org/opensearch/ml/rest/RestMLAgentNameUniquenessIT.java +++ b/plugin/src/test/java/org/opensearch/ml/rest/RestMLAgentNameUniquenessIT.java @@ -97,10 +97,77 @@ public void testDynamicFlip_OnThenOff() throws IOException { assertEquals("Duplicate must be accepted again once flag is flipped off", 200, r2.getStatusLine().getStatusCode()); } + public void testRename_ToUnusedName_Accepted_WhenFlagOn() throws IOException { + updateClusterSettings(SETTING_KEY, true); + + Response r = registerAgent(registerAgentBody("it-agent-rename-src")); + String agentId = (String) parseResponseToMap(r).get("agent_id"); + assertNotNull(agentId); + + Response renamed = updateAgent(agentId, "{\"name\":\"it-agent-rename-dst-" + System.nanoTime() + "\"}"); + assertEquals(200, renamed.getStatusLine().getStatusCode()); + } + + public void testRename_ToExistingName_Rejected_WhenFlagOn() throws IOException { + updateClusterSettings(SETTING_KEY, true); + + // Register two distinct agents + Response first = registerAgent(registerAgentBody("it-agent-rename-existing")); + assertEquals(200, first.getStatusLine().getStatusCode()); + String firstId = (String) parseResponseToMap(first).get("agent_id"); + assertNotNull(firstId); + + Response second = registerAgent(registerAgentBody("it-agent-rename-loser")); + String secondId = (String) parseResponseToMap(second).get("agent_id"); + assertNotNull(secondId); + + // Try to rename the second one onto the first one's name + try { + updateAgent(secondId, "{\"name\":\"it-agent-rename-existing\"}"); + fail("Expected rename to an existing name to be rejected with 409"); + } catch (ResponseException e) { + assertEquals(409, e.getResponse().getStatusLine().getStatusCode()); + String payload = TestHelper.httpEntityToString(e.getResponse().getEntity()); + assertTrue( + "409 response should cite the duplicate name, got: " + payload, + payload.contains("already exists") && payload.contains("it-agent-rename-existing") + ); + assertFalse("409 response must not leak the conflicting agent id, got: " + payload, payload.contains(firstId)); + } + } + + public void testRename_SameName_NoOp_Accepted_WhenFlagOn() throws IOException { + updateClusterSettings(SETTING_KEY, true); + + Response r = registerAgent(registerAgentBody("it-agent-rename-self")); + String agentId = (String) parseResponseToMap(r).get("agent_id"); + assertNotNull(agentId); + + // PUT with the same name: must NOT 409 against itself. + Response same = updateAgent(agentId, "{\"name\":\"it-agent-rename-self\"}"); + assertEquals(200, same.getStatusLine().getStatusCode()); + } + + public void testRename_ToExistingName_Allowed_WhenFlagOff() throws IOException { + // Flag stays off (default from @Before). Rename onto an existing name must succeed (BWC). + Response first = registerAgent(registerAgentBody("it-agent-rename-bwc")); + assertEquals(200, first.getStatusLine().getStatusCode()); + + Response second = registerAgent(registerAgentBody("it-agent-rename-bwc-other")); + String secondId = (String) parseResponseToMap(second).get("agent_id"); + + Response renamed = updateAgent(secondId, "{\"name\":\"it-agent-rename-bwc\"}"); + assertEquals(200, renamed.getStatusLine().getStatusCode()); + } + private static Response registerAgent(String body) throws IOException { return TestHelper.makeRequest(client(), "POST", "/_plugins/_ml/agents/_register", null, TestHelper.toHttpEntity(body), null); } + private static Response updateAgent(String agentId, String body) throws IOException { + return TestHelper.makeRequest(client(), "PUT", "/_plugins/_ml/agents/" + agentId, null, TestHelper.toHttpEntity(body), null); + } + private static String registerAgentBody(String name) { return "{\n" + " \"name\": \"" diff --git a/plugin/src/test/java/org/opensearch/ml/rest/RestMLMemoryContainerNameUniquenessIT.java b/plugin/src/test/java/org/opensearch/ml/rest/RestMLMemoryContainerNameUniquenessIT.java index 485d94bae4..b97b87a38d 100644 --- a/plugin/src/test/java/org/opensearch/ml/rest/RestMLMemoryContainerNameUniquenessIT.java +++ b/plugin/src/test/java/org/opensearch/ml/rest/RestMLMemoryContainerNameUniquenessIT.java @@ -95,11 +95,88 @@ public void testDynamicFlip_OnThenOff() throws IOException { assertEquals("Duplicate must be accepted again once flag is flipped off", 200, r2.getStatusLine().getStatusCode()); } + public void testRename_ToUnusedName_Accepted_WhenFlagOn() throws IOException { + updateClusterSettings(SETTING_KEY, true); + + Response r = createMemoryContainer(createMemoryContainerBody("it-memc-rename-src")); + String id = (String) parseResponseToMap(r).get("memory_container_id"); + assertNotNull(id); + + Response renamed = updateMemoryContainer(id, "{\"name\":\"it-memc-rename-dst-" + System.nanoTime() + "\"}"); + assertEquals(200, renamed.getStatusLine().getStatusCode()); + } + + public void testRename_ToExistingName_Rejected_WhenFlagOn() throws IOException { + updateClusterSettings(SETTING_KEY, true); + + Response first = createMemoryContainer(createMemoryContainerBody("it-memc-rename-existing")); + assertEquals(200, first.getStatusLine().getStatusCode()); + String firstId = (String) parseResponseToMap(first).get("memory_container_id"); + assertNotNull(firstId); + + Response second = createMemoryContainer(createMemoryContainerBody("it-memc-rename-loser")); + String secondId = (String) parseResponseToMap(second).get("memory_container_id"); + assertNotNull(secondId); + + try { + updateMemoryContainer(secondId, "{\"name\":\"it-memc-rename-existing\"}"); + fail("Expected rename to an existing name to be rejected with 409"); + } catch (ResponseException e) { + assertEquals(409, e.getResponse().getStatusLine().getStatusCode()); + String payload = TestHelper.httpEntityToString(e.getResponse().getEntity()); + assertTrue( + "409 response should cite the duplicate name, got: " + payload, + payload.contains("already exists") && payload.contains("it-memc-rename-existing") + ); + assertFalse("409 response must not leak the conflicting memory_container_id, got: " + payload, payload.contains(firstId)); + } + } + + public void testRename_BlankName_NoOp_Accepted_WhenFlagOn() throws IOException { + // A PUT with a blank name must be treated as "no rename" (mirrors the StringUtils.isNotBlank + // gate in TransportUpdateModelGroupAction.updateModelGroup) - must not 409. + updateClusterSettings(SETTING_KEY, true); + + Response r = createMemoryContainer(createMemoryContainerBody("it-memc-rename-blank-src")); + String id = (String) parseResponseToMap(r).get("memory_container_id"); + assertNotNull(id); + + Response blank = updateMemoryContainer(id, "{\"name\":\" \"}"); + assertEquals(200, blank.getStatusLine().getStatusCode()); + } + + public void testRename_SameName_NoOp_Accepted_WhenFlagOn() throws IOException { + updateClusterSettings(SETTING_KEY, true); + + Response r = createMemoryContainer(createMemoryContainerBody("it-memc-rename-self")); + String id = (String) parseResponseToMap(r).get("memory_container_id"); + assertNotNull(id); + + Response same = updateMemoryContainer(id, "{\"name\":\"it-memc-rename-self\"}"); + assertEquals(200, same.getStatusLine().getStatusCode()); + } + + public void testRename_ToExistingName_Allowed_WhenFlagOff() throws IOException { + // Flag stays off (default from @Before). Rename onto an existing name must succeed (BWC). + Response first = createMemoryContainer(createMemoryContainerBody("it-memc-rename-bwc")); + assertEquals(200, first.getStatusLine().getStatusCode()); + + Response second = createMemoryContainer(createMemoryContainerBody("it-memc-rename-bwc-other")); + String secondId = (String) parseResponseToMap(second).get("memory_container_id"); + + Response renamed = updateMemoryContainer(secondId, "{\"name\":\"it-memc-rename-bwc\"}"); + assertEquals(200, renamed.getStatusLine().getStatusCode()); + } + private static Response createMemoryContainer(String body) throws IOException { return TestHelper .makeRequest(client(), "POST", "/_plugins/_ml/memory_containers/_create", null, TestHelper.toHttpEntity(body), null); } + private static Response updateMemoryContainer(String id, String body) throws IOException { + return TestHelper.makeRequest(client(), "PUT", "/_plugins/_ml/memory_containers/" + id, null, TestHelper.toHttpEntity(body), null); + } + private static String createMemoryContainerBody(String name) { return "{\n" + " \"name\": \"" From d4fb1b8b167eacd283ab1f78ce8913ad1ba9b8ae Mon Sep 17 00:00:00 2001 From: rithin-pullela-aws Date: Tue, 5 May 2026 11:00:31 -0700 Subject: [PATCH 3/3] fix: Reject blank memory-container names and tighten update-path guards - MLCreateMemoryContainerInput: reject null *or* blank name (was null-only), matching MLAgent's constructor-level check. Prevents the create path from accepting a whitespace-only name that would later look valid in search. - MLUpdateMemoryContainerInput: reject a non-null blank name at parse time. Null still means "no rename" (update's name field is optional), but a caller sending " " now fails fast with 400 instead of silently either 409'ing or overwriting the stored name with whitespace in continueUpdate. - TransportUpdateMemoryContainerAction.continueUpdate: tighten the guard from (newName != null) to StringUtils.isNotBlank(newName) as defense-in-depth behind the input-layer check. - UpdateAgentTransportAction: scope the uniqueness search by the retrieved agent's tenantId instead of the update input's, making the authoritative tenant source consistent with the memory-container path. - Standardize on org.apache.commons.lang3.StringUtils across the three name-blank checks so the predicate reads the same everywhere. - IT testRename_BlankName: was asserting 200 (a false green - it didn't read the doc back), now asserts 400 + reads the name back to confirm it wasn't mutated. testRename_ToUnusedName also reads back to confirm the rename actually persisted. - UTs for blank/empty name rejection on both inputs, plus a positive test that null name is still allowed on update. Signed-off-by: rithin-pullela-aws --- .../MLCreateMemoryContainerInput.java | 5 ++-- .../memory/MLUpdateMemoryContainerInput.java | 5 ++++ .../MLCreateMemoryContainerInputTests.java | 10 +++++++ .../MLUpdateMemoryContainerInputTests.java | 17 ++++++++++++ .../agents/UpdateAgentTransportAction.java | 2 +- .../TransportUpdateMemoryContainerAction.java | 2 +- ...RestMLMemoryContainerNameUniquenessIT.java | 27 ++++++++++++++----- 7 files changed, 58 insertions(+), 10 deletions(-) diff --git a/common/src/main/java/org/opensearch/ml/common/transport/memorycontainer/MLCreateMemoryContainerInput.java b/common/src/main/java/org/opensearch/ml/common/transport/memorycontainer/MLCreateMemoryContainerInput.java index d73c8c132b..2708b11b36 100644 --- a/common/src/main/java/org/opensearch/ml/common/transport/memorycontainer/MLCreateMemoryContainerInput.java +++ b/common/src/main/java/org/opensearch/ml/common/transport/memorycontainer/MLCreateMemoryContainerInput.java @@ -17,6 +17,7 @@ import java.util.List; import java.util.regex.Pattern; +import org.apache.commons.lang3.StringUtils; import org.opensearch.OpenSearchParseException; import org.opensearch.core.common.io.stream.StreamInput; import org.opensearch.core.common.io.stream.StreamOutput; @@ -52,8 +53,8 @@ public MLCreateMemoryContainerInput( String tenantId, List backendRoles ) { - if (name == null) { - throw new IllegalArgumentException("name is null"); + if (StringUtils.isBlank(name)) { + throw new IllegalArgumentException("name cannot be null or blank"); } this.name = name; this.description = description; diff --git a/common/src/main/java/org/opensearch/ml/common/transport/memorycontainer/memory/MLUpdateMemoryContainerInput.java b/common/src/main/java/org/opensearch/ml/common/transport/memorycontainer/memory/MLUpdateMemoryContainerInput.java index 77d7ed0114..7cd53d421d 100644 --- a/common/src/main/java/org/opensearch/ml/common/transport/memorycontainer/memory/MLUpdateMemoryContainerInput.java +++ b/common/src/main/java/org/opensearch/ml/common/transport/memorycontainer/memory/MLUpdateMemoryContainerInput.java @@ -15,6 +15,7 @@ import java.util.ArrayList; import java.util.List; +import org.apache.commons.lang3.StringUtils; import org.opensearch.core.common.io.stream.StreamInput; import org.opensearch.core.common.io.stream.StreamOutput; import org.opensearch.core.common.io.stream.Writeable; @@ -39,6 +40,10 @@ public class MLUpdateMemoryContainerInput implements ToXContentObject, Writeable @Builder public MLUpdateMemoryContainerInput(String name, String description, List backendRoles, MemoryConfiguration configuration) { + // Null = "no rename" (name is optional on update), but a non-null blank string is always wrong. + if (name != null && StringUtils.isBlank(name)) { + throw new IllegalArgumentException("name cannot be blank"); + } this.name = name; this.description = description; validateBackendRoles(backendRoles); diff --git a/common/src/test/java/org/opensearch/ml/common/transport/memorycontainer/MLCreateMemoryContainerInputTests.java b/common/src/test/java/org/opensearch/ml/common/transport/memorycontainer/MLCreateMemoryContainerInputTests.java index 5917d837bd..874e292e79 100644 --- a/common/src/test/java/org/opensearch/ml/common/transport/memorycontainer/MLCreateMemoryContainerInputTests.java +++ b/common/src/test/java/org/opensearch/ml/common/transport/memorycontainer/MLCreateMemoryContainerInputTests.java @@ -131,6 +131,16 @@ public void testConstructorWithNullNameDirectConstructor() { .build(); } + @Test(expected = IllegalArgumentException.class) + public void testConstructorWithBlankName() { + MLCreateMemoryContainerInput.builder().name(" ").description("blank name").build(); + } + + @Test(expected = IllegalArgumentException.class) + public void testConstructorWithEmptyName() { + MLCreateMemoryContainerInput.builder().name("").description("empty name").build(); + } + @Test public void testStreamInputOutput() throws IOException { BytesStreamOutput bytesStreamOutput = new BytesStreamOutput(); diff --git a/common/src/test/java/org/opensearch/ml/common/transport/memorycontainer/memory/MLUpdateMemoryContainerInputTests.java b/common/src/test/java/org/opensearch/ml/common/transport/memorycontainer/memory/MLUpdateMemoryContainerInputTests.java index 927527c003..abaa61fad4 100644 --- a/common/src/test/java/org/opensearch/ml/common/transport/memorycontainer/memory/MLUpdateMemoryContainerInputTests.java +++ b/common/src/test/java/org/opensearch/ml/common/transport/memorycontainer/memory/MLUpdateMemoryContainerInputTests.java @@ -33,6 +33,23 @@ public class MLUpdateMemoryContainerInputTests { + @Test(expected = IllegalArgumentException.class) + public void testConstructorRejectsBlankName() { + MLUpdateMemoryContainerInput.builder().name(" ").build(); + } + + @Test(expected = IllegalArgumentException.class) + public void testConstructorRejectsEmptyName() { + MLUpdateMemoryContainerInput.builder().name("").build(); + } + + @Test + public void testConstructorAllowsNullName() { + // Null name on update means "no rename" - must be allowed. + MLUpdateMemoryContainerInput input = MLUpdateMemoryContainerInput.builder().name(null).description("desc").build(); + assertNull(input.getName()); + } + @Test public void testConstructor() { List backendRoles = Arrays.asList("role1", "role2"); diff --git a/plugin/src/main/java/org/opensearch/ml/action/agents/UpdateAgentTransportAction.java b/plugin/src/main/java/org/opensearch/ml/action/agents/UpdateAgentTransportAction.java index 758f9f1538..b1ff2615f1 100644 --- a/plugin/src/main/java/org/opensearch/ml/action/agents/UpdateAgentTransportAction.java +++ b/plugin/src/main/java/org/opensearch/ml/action/agents/UpdateAgentTransportAction.java @@ -218,7 +218,7 @@ private void validateAgentNameUniquenessForUpdate( } NameUniquenessHelper - .searchByExactName(client, sdkClient, ML_AGENT_INDEX, newName, updateInput.getTenantId(), ActionListener.wrap(response -> { + .searchByExactName(client, sdkClient, ML_AGENT_INDEX, newName, originalAgent.getTenantId(), ActionListener.wrap(response -> { if (response == null) { listener.onResponse(null); return; diff --git a/plugin/src/main/java/org/opensearch/ml/action/memorycontainer/TransportUpdateMemoryContainerAction.java b/plugin/src/main/java/org/opensearch/ml/action/memorycontainer/TransportUpdateMemoryContainerAction.java index 52efd06f5c..e9bed566cf 100644 --- a/plugin/src/main/java/org/opensearch/ml/action/memorycontainer/TransportUpdateMemoryContainerAction.java +++ b/plugin/src/main/java/org/opensearch/ml/action/memorycontainer/TransportUpdateMemoryContainerAction.java @@ -154,7 +154,7 @@ private void continueUpdate( try { // Prepare the update Map updateFields = new HashMap<>(); - if (newName != null) { + if (StringUtils.isNotBlank(newName)) { updateFields.put(NAME_FIELD, newName); } if (newDescription != null) { diff --git a/plugin/src/test/java/org/opensearch/ml/rest/RestMLMemoryContainerNameUniquenessIT.java b/plugin/src/test/java/org/opensearch/ml/rest/RestMLMemoryContainerNameUniquenessIT.java index b97b87a38d..3da3a9f4b3 100644 --- a/plugin/src/test/java/org/opensearch/ml/rest/RestMLMemoryContainerNameUniquenessIT.java +++ b/plugin/src/test/java/org/opensearch/ml/rest/RestMLMemoryContainerNameUniquenessIT.java @@ -102,8 +102,10 @@ public void testRename_ToUnusedName_Accepted_WhenFlagOn() throws IOException { String id = (String) parseResponseToMap(r).get("memory_container_id"); assertNotNull(id); - Response renamed = updateMemoryContainer(id, "{\"name\":\"it-memc-rename-dst-" + System.nanoTime() + "\"}"); + String newName = "it-memc-rename-dst-" + System.nanoTime(); + Response renamed = updateMemoryContainer(id, "{\"name\":\"" + newName + "\"}"); assertEquals(200, renamed.getStatusLine().getStatusCode()); + assertEquals("Rename should have been persisted", newName, getMemoryContainerName(id)); } public void testRename_ToExistingName_Rejected_WhenFlagOn() throws IOException { @@ -132,17 +134,24 @@ public void testRename_ToExistingName_Rejected_WhenFlagOn() throws IOException { } } - public void testRename_BlankName_NoOp_Accepted_WhenFlagOn() throws IOException { - // A PUT with a blank name must be treated as "no rename" (mirrors the StringUtils.isNotBlank - // gate in TransportUpdateModelGroupAction.updateModelGroup) - must not 409. + public void testRename_BlankName_Rejected_WhenFlagOn() throws IOException { + // A PUT with a blank name must be rejected up front (mirrors MLAgentUpdateInput.validate()): + // we'd rather 400 than silently either 409 or overwrite the stored name with whitespace. updateClusterSettings(SETTING_KEY, true); Response r = createMemoryContainer(createMemoryContainerBody("it-memc-rename-blank-src")); String id = (String) parseResponseToMap(r).get("memory_container_id"); assertNotNull(id); - Response blank = updateMemoryContainer(id, "{\"name\":\" \"}"); - assertEquals(200, blank.getStatusLine().getStatusCode()); + try { + updateMemoryContainer(id, "{\"name\":\" \"}"); + fail("Expected blank name on rename to be rejected"); + } catch (ResponseException e) { + assertEquals(400, e.getResponse().getStatusLine().getStatusCode()); + } + + // Confirm the stored name was not mutated. + assertEquals("it-memc-rename-blank-src", getMemoryContainerName(id)); } public void testRename_SameName_NoOp_Accepted_WhenFlagOn() throws IOException { @@ -177,6 +186,12 @@ private static Response updateMemoryContainer(String id, String body) throws IOE return TestHelper.makeRequest(client(), "PUT", "/_plugins/_ml/memory_containers/" + id, null, TestHelper.toHttpEntity(body), null); } + private String getMemoryContainerName(String id) throws IOException { + Response r = TestHelper.makeRequest(client(), "GET", "/_plugins/_ml/memory_containers/" + id, null, "", null); + assertEquals(200, r.getStatusLine().getStatusCode()); + return (String) parseResponseToMap(r).get("name"); + } + private static String createMemoryContainerBody(String name) { return "{\n" + " \"name\": \""