Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,7 @@
import static org.apache.solr.mcp.server.collection.CollectionUtils.getInteger;
import static org.apache.solr.mcp.server.collection.CollectionUtils.getLong;
import static org.apache.solr.mcp.server.util.JsonUtils.toJson;
import static org.apache.solr.mcp.server.util.ToolArguments.requireCollection;

import com.fasterxml.jackson.databind.ObjectMapper;
import io.micrometer.observation.annotation.Observed;
Expand Down Expand Up @@ -259,9 +260,6 @@ public class CollectionService {
/** Default replication factor for new collections */
private static final int DEFAULT_REPLICATION_FACTOR = 1;

/** Error message for blank collection name validation */
private static final String BLANK_COLLECTION_NAME_ERROR = "Collection name must not be blank";

/** SolrJ client for communicating with Solr server */
private final SolrClient solrClient;

Expand Down Expand Up @@ -519,6 +517,8 @@ public List<String> listCollections() throws SolrServerException, IOException {
public SolrMetrics getCollectionStats(
@McpToolParam(description = "Solr collection to get stats/metrics for") String collection)
throws SolrServerException, IOException {
requireCollection(collection);

// Extract actual collection name from shard name if needed
String actualCollection = extractCollectionName(collection);

Expand Down Expand Up @@ -1067,6 +1067,8 @@ private boolean validateCollectionExists(String collection) throws SolrServerExc
annotations = @McpTool.McpAnnotations(readOnlyHint = true),
description = "Check health of a Solr collection")
public SolrHealthStatus checkHealth(@McpToolParam(description = "Solr collection") String collection) {
requireCollection(collection);

String actualCollection = extractCollectionName(collection);
try {
// Ping Solr
Expand Down Expand Up @@ -1135,9 +1137,7 @@ public CollectionCreationResult createCollection(
required = false) @Nullable Integer replicationFactor)
throws SolrServerException, IOException {

if (name == null || name.isBlank()) {
throw new IllegalArgumentException(BLANK_COLLECTION_NAME_ERROR);
}
requireCollection(name);

String effectiveConfigSet = configSet != null ? configSet : DEFAULT_CONFIGSET;
int effectiveShards = numShards != null ? numShards : DEFAULT_NUM_SHARDS;
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,8 @@
*/
package org.apache.solr.mcp.server.indexing;

import static org.apache.solr.mcp.server.util.ToolArguments.requireCollection;

import io.micrometer.observation.annotation.Observed;
import java.io.IOException;
import java.util.List;
Expand Down Expand Up @@ -213,6 +215,8 @@ public IndexingService(SolrClient solrClient, IndexingDocumentCreator indexingDo
public String indexJsonDocuments(@McpToolParam(description = "Solr collection to index into") String collection,
@McpToolParam(description = "JSON string containing documents to index") String json)
throws IOException, SolrServerException {
requireCollection(collection);

List<SolrInputDocument> schemalessDoc = indexingDocumentCreator.createSchemalessDocumentsFromJson(json);
int successCount = indexDocuments(collection, schemalessDoc);
return "Successfully indexed " + successCount + " of " + schemalessDoc.size() + " documents into collection '"
Expand Down Expand Up @@ -288,6 +292,8 @@ public String indexJsonDocuments(@McpToolParam(description = "Solr collection to
public String indexCsvDocuments(@McpToolParam(description = "Solr collection to index into") String collection,
@McpToolParam(description = "CSV string containing documents to index") String csv)
throws IOException, SolrServerException {
requireCollection(collection);

List<SolrInputDocument> schemalessDoc = indexingDocumentCreator.createSchemalessDocumentsFromCsv(csv);
int successCount = indexDocuments(collection, schemalessDoc);
return "Successfully indexed " + successCount + " of " + schemalessDoc.size() + " documents into collection '"
Expand Down Expand Up @@ -387,6 +393,8 @@ public String indexCsvDocuments(@McpToolParam(description = "Solr collection to
public String indexXmlDocuments(@McpToolParam(description = "Solr collection to index into") String collection,
@McpToolParam(description = "XML string containing documents to index") String xml)
throws ParserConfigurationException, SAXException, IOException, SolrServerException {
requireCollection(collection);

List<SolrInputDocument> schemalessDoc = indexingDocumentCreator.createSchemalessDocumentsFromXml(xml);
int successCount = indexDocuments(collection, schemalessDoc);
return "Successfully indexed " + successCount + " of " + schemalessDoc.size() + " documents into collection '"
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,7 @@

import static org.apache.solr.mcp.server.util.JsonUtils.toJson;
import static org.apache.solr.mcp.server.util.PromptText.optionalCodeBlock;
import static org.apache.solr.mcp.server.util.ToolArguments.requireCollection;

import com.fasterxml.jackson.databind.ObjectMapper;
import io.micrometer.observation.annotation.Observed;
Expand Down Expand Up @@ -182,6 +183,8 @@ public SchemaService(SolrClient solrClient, ObjectMapper objectMapper) {
description = "Schema definition for a Solr collection including fields, field types, and copy fields",
mimeType = "application/json")
public String getSchemaResource(String collection) {
requireCollection(collection);

try {
return toJson(objectMapper, getSchema(collection));
} catch (Exception e) {
Expand Down Expand Up @@ -277,6 +280,8 @@ public String getSchemaResource(String collection) {
annotations = @McpTool.McpAnnotations(readOnlyHint = true),
description = "Get schema for a Solr collection")
public SchemaRepresentation getSchema(String collection) throws Exception {
requireCollection(collection);

SchemaRequest schemaRequest = new SchemaRequest();
return schemaRequest.process(solrClient, collection).getSchemaRepresentation();
}
Expand Down Expand Up @@ -482,15 +487,9 @@ private AnalyzerDefinition toAnalyzerDefinition(Object raw) {
return def;
}

private static void requireCollection(String collection) {
if (collection == null || collection.isBlank()) {
throw new IllegalArgumentException("Collection name must not be blank");
}
}

private static void requireNonEmpty(List<?> list, String name) {
if (list == null || list.isEmpty()) {
throw new IllegalArgumentException(name + " must not be empty");
throw new IllegalArgumentException(name + " cannot be null or empty");
}
}

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,8 @@
*/
package org.apache.solr.mcp.server.search;

import static org.apache.solr.mcp.server.util.ToolArguments.requireCollection;

import io.micrometer.observation.annotation.Observed;
import java.io.IOException;
import java.util.ArrayList;
Expand Down Expand Up @@ -309,6 +311,8 @@ public SearchResponse search(@McpToolParam(description = "Solr collection to que
@McpToolParam(description = "Number of rows to return", required = false) @Nullable Integer rows)
throws SolrServerException, IOException {

requireCollection(collection);

// query
final SolrQuery solrQuery = new SolrQuery("*:*");
if (StringUtils.hasText(query)) {
Expand Down
62 changes: 62 additions & 0 deletions src/main/java/org/apache/solr/mcp/server/util/ToolArguments.java
Original file line number Diff line number Diff line change
@@ -0,0 +1,62 @@
/*
* Licensed to the Apache Software Foundation (ASF) under one or more
* contributor license agreements. See the NOTICE file distributed with
* this work for additional information regarding copyright ownership.
* The ASF licenses this file to You under the Apache License, Version 2.0
* (the "License"); you may not use this file except in compliance with
* the License. You may obtain a copy of the License at
*
* http://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
package org.apache.solr.mcp.server.util;

/**
* Validation of arguments arriving at MCP tool boundaries.
*
* <p>
* Every {@code @McpTool} method that takes a collection name calls
* {@link #requireCollection(String)} first, so a client that omits the argument
* or sends an empty string gets one identical, actionable message instead of a
* per-service variant or an opaque downstream failure.
*
* <p>
* <strong>Why a runtime check in null-marked code:</strong> the server is
* {@code @NullMarked} (see {@code org.apache.solr.mcp.server.package-info}) and
* NullAway runs as a build error, but that analysis only binds callers the
* compiler can see. MCP tool methods are invoked reflectively by the Spring AI
* annotation runtime, which resolves each parameter with a plain lookup against
* the request's argument map and passes the result straight through — a missing
* or null JSON value therefore reaches the method as {@code null} no matter
* what the annotations declare. {@code @McpToolParam(required = true)} only
* marks the parameter required in the advertised JSON schema; the server does
* not validate incoming arguments against it. These checks are the trust
* boundary, not redundant defensive coding.
*/
public final class ToolArguments {

/** Message used whenever a collection name is missing, empty, or blank. */
public static final String BLANK_COLLECTION_NAME_ERROR = "Collection name cannot be null or empty";

private ToolArguments() {
}

/**
* Rejects a missing, empty, or whitespace-only collection name.
*
* @param collection
* the collection name supplied by the MCP client
* @throws IllegalArgumentException
* if {@code collection} is null or contains only whitespace
*/
public static void requireCollection(String collection) {
if (collection == null || collection.isBlank()) {
throw new IllegalArgumentException(BLANK_COLLECTION_NAME_ERROR);
}
}
}
Original file line number Diff line number Diff line change
@@ -0,0 +1,134 @@
/*
* Licensed to the Apache Software Foundation (ASF) under one or more
* contributor license agreements. See the NOTICE file distributed with
* this work for additional information regarding copyright ownership.
* The ASF licenses this file to You under the Apache License, Version 2.0
* (the "License"); you may not use this file except in compliance with
* the License. You may obtain a copy of the License at
*
* http://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
package org.apache.solr.mcp.server;

import static org.apache.solr.mcp.server.util.ToolArguments.BLANK_COLLECTION_NAME_ERROR;
import static org.assertj.core.api.Assertions.assertThatThrownBy;

import com.fasterxml.jackson.databind.ObjectMapper;
import org.apache.solr.client.solrj.SolrClient;
import org.apache.solr.mcp.server.collection.CollectionService;
import org.apache.solr.mcp.server.indexing.IndexingService;
import org.apache.solr.mcp.server.indexing.documentcreator.IndexingDocumentCreator;
import org.apache.solr.mcp.server.schema.SchemaService;
import org.apache.solr.mcp.server.search.SearchService;
import org.assertj.core.api.ThrowableAssert.ThrowingCallable;
import org.junit.jupiter.api.BeforeEach;
import org.junit.jupiter.api.Test;
import org.junit.jupiter.api.condition.DisabledInNativeImage;
import org.junit.jupiter.api.extension.ExtendWith;
import org.mockito.Mock;
import org.mockito.junit.jupiter.MockitoExtension;

/**
* Verifies that every MCP tool method taking a collection name rejects a
* missing, empty, or whitespace-only value with the same message.
*
* <p>
* The services are exercised directly with mocked collaborators: validation
* must happen before any Solr call, so a client that omits the argument gets an
* actionable message rather than an opaque downstream failure. Asserting
* against {@code BLANK_COLLECTION_NAME_ERROR} rather than a copied string
* literal is deliberate — it is what keeps the four services from drifting
* apart.
*/
@ExtendWith(MockitoExtension.class)
@DisabledInNativeImage
class CollectionNameValidationTest {

private static final String BLANK = " ";

@Mock
private SolrClient solrClient;

@Mock
private IndexingDocumentCreator indexingDocumentCreator;

private CollectionService collectionService;
private IndexingService indexingService;
private SchemaService schemaService;
private SearchService searchService;

@BeforeEach
void setUp() {
ObjectMapper objectMapper = new ObjectMapper();
collectionService = new CollectionService(solrClient, objectMapper);
indexingService = new IndexingService(solrClient, indexingDocumentCreator);
schemaService = new SchemaService(solrClient, objectMapper);
searchService = new SearchService(solrClient);
}

private static void assertRejectsBlankCollection(ThrowingCallable withNull, ThrowingCallable withBlank) {
assertThatThrownBy(withNull).isInstanceOf(IllegalArgumentException.class)
.hasMessage(BLANK_COLLECTION_NAME_ERROR);
assertThatThrownBy(withBlank).isInstanceOf(IllegalArgumentException.class)
.hasMessage(BLANK_COLLECTION_NAME_ERROR);
}

@Test
void createCollectionRejectsBlankCollectionName() {
assertRejectsBlankCollection(() -> collectionService.createCollection(null, null, null, null),
() -> collectionService.createCollection(BLANK, null, null, null));
}

@Test
void getCollectionStatsRejectsBlankCollectionName() {
assertRejectsBlankCollection(() -> collectionService.getCollectionStats(null),
() -> collectionService.getCollectionStats(BLANK));
}

@Test
void checkHealthRejectsBlankCollectionName() {
assertRejectsBlankCollection(() -> collectionService.checkHealth(null),
() -> collectionService.checkHealth(BLANK));
}

@Test
void indexJsonDocumentsRejectsBlankCollectionName() {
assertRejectsBlankCollection(() -> indexingService.indexJsonDocuments(null, "[]"),
() -> indexingService.indexJsonDocuments(BLANK, "[]"));
}

@Test
void indexCsvDocumentsRejectsBlankCollectionName() {
assertRejectsBlankCollection(() -> indexingService.indexCsvDocuments(null, "id\n1"),
() -> indexingService.indexCsvDocuments(BLANK, "id\n1"));
}

@Test
void indexXmlDocumentsRejectsBlankCollectionName() {
assertRejectsBlankCollection(() -> indexingService.indexXmlDocuments(null, "<docs/>"),
() -> indexingService.indexXmlDocuments(BLANK, "<docs/>"));
}

@Test
void getSchemaRejectsBlankCollectionName() {
assertRejectsBlankCollection(() -> schemaService.getSchema(null), () -> schemaService.getSchema(BLANK));
}

@Test
void getSchemaResourceRejectsBlankCollectionName() {
assertRejectsBlankCollection(() -> schemaService.getSchemaResource(null),
() -> schemaService.getSchemaResource(BLANK));
}

@Test
void searchRejectsBlankCollectionName() {
assertRejectsBlankCollection(() -> searchService.search(null, null, null, null, null, null, null),
() -> searchService.search(BLANK, null, null, null, null, null, null));
}
}