From df3b99bb48581783621e9ed4844c694c8bd89933 Mon Sep 17 00:00:00 2001 From: adj5672 Date: Mon, 14 Sep 2026 14:08:04 +0900 Subject: [PATCH] Add OpenAI Tool Search Index Add an OpenAI-backed ToolIndex implementation that uses the Responses API tool-search capability to select matching tools from session-scoped tool references. Add OpenAI SDK core as an optional dependency and cover builder validation, request creation, session isolation, index clearing, metadata extraction, and client closing behavior with unit tests. Signed-off-by: adj5672 --- spring-ai-tool-search-tool/pom.xml | 6 + .../index/openai/OpenAiToolIndex.java | 252 ++++++++++++++++ .../toolsearch/index/openai/package-info.java | 4 + .../index/openai/OpenAiToolIndexTests.java | 272 ++++++++++++++++++ 4 files changed, 534 insertions(+) create mode 100644 spring-ai-tool-search-tool/src/main/java/org/springframework/ai/tool/toolsearch/index/openai/OpenAiToolIndex.java create mode 100644 spring-ai-tool-search-tool/src/main/java/org/springframework/ai/tool/toolsearch/index/openai/package-info.java create mode 100644 spring-ai-tool-search-tool/src/test/java/org/springframework/ai/tool/toolsearch/index/openai/OpenAiToolIndexTests.java diff --git a/spring-ai-tool-search-tool/pom.xml b/spring-ai-tool-search-tool/pom.xml index f58d95edce..7c06d33b57 100644 --- a/spring-ai-tool-search-tool/pom.xml +++ b/spring-ai-tool-search-tool/pom.xml @@ -35,6 +35,12 @@ true + + com.openai + openai-java-core + ${openai-sdk.version} + true + diff --git a/spring-ai-tool-search-tool/src/main/java/org/springframework/ai/tool/toolsearch/index/openai/OpenAiToolIndex.java b/spring-ai-tool-search-tool/src/main/java/org/springframework/ai/tool/toolsearch/index/openai/OpenAiToolIndex.java new file mode 100644 index 0000000000..1175382fc9 --- /dev/null +++ b/spring-ai-tool-search-tool/src/main/java/org/springframework/ai/tool/toolsearch/index/openai/OpenAiToolIndex.java @@ -0,0 +1,252 @@ +package org.springframework.ai.tool.toolsearch.index.openai; + +import com.openai.client.OpenAIClient; +import com.openai.errors.OpenAIException; +import com.openai.models.ChatModel; +import com.openai.models.responses.FunctionTool; +import com.openai.models.responses.Response; +import com.openai.models.responses.ResponseCreateParams; +import com.openai.models.responses.ResponseOutputItem; +import com.openai.models.responses.ResponseToolSearchOutputItem; +import com.openai.models.responses.Tool; +import com.openai.models.responses.ToolSearchTool; +import org.apache.commons.logging.Log; +import org.apache.commons.logging.LogFactory; +import org.jspecify.annotations.Nullable; +import org.springframework.ai.tool.toolsearch.ToolIndex; +import org.springframework.ai.tool.toolsearch.ToolReference; +import org.springframework.ai.tool.toolsearch.ToolSearchRequest; +import org.springframework.ai.tool.toolsearch.ToolSearchResponse; +import org.springframework.ai.tool.toolsearch.ToolSearchResponse.SearchMetadata; +import org.springframework.util.Assert; + +import java.io.Closeable; +import java.util.ArrayList; +import java.util.List; +import java.util.Map; +import java.util.Optional; +import java.util.concurrent.ConcurrentHashMap; + +/** + * OpenAI-based tool searcher that delegates tool matching to the OpenAI Responses API + * tool-search capability. + *

+ * Tools are stored per session and sent as function tool definitions when a search is + * executed. The OpenAI model returns the function tools that best match the search query. + *

+ * OpenAI + * Developers - Tool Search + * + */ +public class OpenAiToolIndex implements ToolIndex, Closeable { + + private static final Log logger = LogFactory.getLog(OpenAiToolIndex.class); + + private static final ToolSearchTool TOOL_SEARCH_TOOL = ToolSearchTool.builder() + .execution(ToolSearchTool.Execution.SERVER) + .build(); + + private final OpenAIClient openAiClient; + + private final String model; + + /** + * Function tools indexed by session. OpenAI's tool-search API accepts function tool + * definitions, so non-function tool references are represented by their name and + * description only. + */ + private final Map> sessionTools = new ConcurrentHashMap<>(); + + private OpenAiToolIndex(OpenAIClient openAiClient, String model) { + this.openAiClient = openAiClient; + this.model = model; + } + + @Override + public void indexTool(String sessionId, ToolReference toolReference) { + this.indexTools(sessionId, List.of(toolReference)); + } + + @Override + public void indexTools(String sessionId, List toolReferences) { + List sessionTools = this.sessionTools.getOrDefault(sessionId, List.of()); + List newSessionTools = new ArrayList<>(sessionTools); + + toolReferences.stream().map(this::toOpenAiFunctionTool).forEach(newSessionTools::add); + + this.sessionTools.put(sessionId, List.copyOf(newSessionTools)); + + } + + @Override + public ToolSearchResponse search(ToolSearchRequest toolSearchRequest) { + String query = toolSearchRequest.query(); + String sessionId = toolSearchRequest.sessionId(); + + List sessionTools = this.sessionTools.getOrDefault(sessionId, List.of()); + if (sessionTools.isEmpty()) { + return ToolSearchResponse.builder().build(); + } + + Response response = createToolSearchResponse(query, sessionTools); + + List toolReferences = getToolSearchOutputTools(response).stream() + .map(this::toToolReference) + .toList(); + + SearchMetadata searchMetadata = buildSearchMetadata(query, response); + + return ToolSearchResponse.builder() + .toolReferences(toolReferences) + .totalMatches(toolReferences.size()) + .searchMetadata(searchMetadata) + .build(); + } + + @Override + public void clearIndex(String sessionId) { + this.sessionTools.remove(sessionId); + } + + /** + * Creates a Responses API request that asks OpenAI to select matching tools for the + * query. + * @param query the search query + * @param tools the function tools available for the current session + * @return the OpenAI response containing tool-search output items + */ + private Response createToolSearchResponse(String query, List tools) { + ResponseCreateParams.Builder paramsBuilder = ResponseCreateParams.builder() + .model(this.model) + .input(query) + .parallelToolCalls(false) + .addTool(TOOL_SEARCH_TOOL); + tools.forEach(paramsBuilder::addTool); + + return this.openAiClient.responses().create(paramsBuilder.build()); + } + + /** + * Extracts function tools selected by OpenAI's tool-search output. + * @param response the OpenAI response to inspect + * @return selected function tools from tool-search output items + */ + private List getToolSearchOutputTools(Response response) { + List toolSearchOutputItems = response.output() + .stream() + .filter(ResponseOutputItem::isToolSearchOutput) + .map(ResponseOutputItem::asToolSearchOutput) + .toList(); + + return toolSearchOutputItems.stream() + .flatMap(item -> item.tools().stream()) + .filter(Tool::isFunction) + .map(Tool::asFunction) + .toList(); + } + + /** + * Builds search metadata from the OpenAI response timestamps. + * @param query the original search query + * @param response the OpenAI response + * @return metadata describing the search execution + */ + private SearchMetadata buildSearchMetadata(String query, Response response) { + SearchMetadata.Builder builder = SearchMetadata.builder() + .searchType(this.getClass().getSimpleName()) + .query(query); + + response.completedAt() + .map(completedAt -> completedAt * 1_000 - response.createdAt() * 1_000) + .map(Double::longValue) + .ifPresent(builder::searchTimeMs); + + return builder.build(); + } + + /** + * Converts a {@link ToolReference} into an OpenAI function tool definition. + * @param toolReference the indexed tool reference + * @return an OpenAI function tool for tool-search requests + */ + private FunctionTool toOpenAiFunctionTool(ToolReference toolReference) { + return FunctionTool.builder() + .name(toolReference.toolName()) + .description(toolReference.summary()) + // ToolReference contains searchable metadata, but not an input schema. + .parameters(Optional.empty()) + // Strict tool validation is configured by ToolCallingOptions. + .strict(false) + .build(); + } + + /** + * Converts an OpenAI function tool returned by tool search into a + * {@link ToolReference}. + * @param functionTool the selected OpenAI function tool + * @return a tool reference for the search response + */ + private ToolReference toToolReference(FunctionTool functionTool) { + return ToolReference.builder() + .toolName(functionTool.name()) + .summary(functionTool.description().orElse("")) + .build(); + } + + public static Builder builder() { + return new Builder(); + } + + @Override + public void close() { + this.openAiClient.close(); + } + + /** + * Builder for {@link OpenAiToolIndex}. + */ + public static class Builder { + + @Nullable private OpenAIClient openAiClient; + + @Nullable private String model; + + /** + * Configure the OpenAI client used for tool-search requests. + * @param openAiClient the OpenAI client + * @return this builder + */ + public Builder openAiClient(OpenAIClient openAiClient) { + this.openAiClient = openAiClient; + return this; + } + + /** + * Configure the OpenAI model used for tool-search requests. + * @param model the OpenAI model name + * @return this builder + */ + public Builder model(String model) { + this.model = model; + return this; + } + + /** + * Build an {@link OpenAiToolIndex}. + * @return a configured OpenAI-backed tool index + */ + public OpenAiToolIndex build() { + Assert.notNull(this.openAiClient, "openAiClient must not be null"); + Assert.notNull(this.model, "model must not be null"); + + ChatModel chatModel = ChatModel.of(this.model); + if (chatModel.value().ordinal() > ChatModel.GPT_5_4_MINI.value().ordinal()) { + throw new OpenAIException("Only gpt-5.4 and later models support tool search."); + } + + return new OpenAiToolIndex(this.openAiClient, this.model); + } + + } + +} diff --git a/spring-ai-tool-search-tool/src/main/java/org/springframework/ai/tool/toolsearch/index/openai/package-info.java b/spring-ai-tool-search-tool/src/main/java/org/springframework/ai/tool/toolsearch/index/openai/package-info.java new file mode 100644 index 0000000000..98955abf1f --- /dev/null +++ b/spring-ai-tool-search-tool/src/main/java/org/springframework/ai/tool/toolsearch/index/openai/package-info.java @@ -0,0 +1,4 @@ +@NullMarked +package org.springframework.ai.tool.toolsearch.index.openai; + +import org.jspecify.annotations.NullMarked; diff --git a/spring-ai-tool-search-tool/src/test/java/org/springframework/ai/tool/toolsearch/index/openai/OpenAiToolIndexTests.java b/spring-ai-tool-search-tool/src/test/java/org/springframework/ai/tool/toolsearch/index/openai/OpenAiToolIndexTests.java new file mode 100644 index 0000000000..175de93eb0 --- /dev/null +++ b/spring-ai-tool-search-tool/src/test/java/org/springframework/ai/tool/toolsearch/index/openai/OpenAiToolIndexTests.java @@ -0,0 +1,272 @@ +/* + * Copyright 2023-present the original author or authors. + * + * Licensed 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 + * + * https://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.springframework.ai.tool.toolsearch.index.openai; + +import java.util.List; +import java.util.Objects; +import java.util.Optional; +import java.util.stream.Stream; + +import com.openai.client.OpenAIClient; +import com.openai.errors.OpenAIException; +import com.openai.models.ChatModel; +import com.openai.models.responses.FunctionTool; +import com.openai.models.responses.Response; +import com.openai.models.responses.Response.ToolChoice; +import com.openai.models.responses.ResponseCreateParams; +import com.openai.models.responses.ResponseOutputItem; +import com.openai.models.responses.ResponseToolSearchOutputItem; +import com.openai.models.responses.Tool; +import com.openai.models.responses.ToolChoiceOptions; +import com.openai.services.blocking.ResponseService; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.ArgumentCaptor; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; + +import org.springframework.ai.tool.toolsearch.ToolReference; +import org.springframework.ai.tool.toolsearch.ToolSearchRequest; +import org.springframework.ai.tool.toolsearch.ToolSearchResponse; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatIllegalArgumentException; +import static org.assertj.core.api.Assertions.assertThatThrownBy; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.Mockito.never; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; + +/** + * Unit tests for {@link OpenAiToolIndex}. + * + * @author Christian Tzolov + */ +@ExtendWith(MockitoExtension.class) +class OpenAiToolIndexTests { + + private static final String SESSION = "session-1"; + + private static final String MODEL = ChatModel.GPT_5_4_MINI.asString(); + + @Mock + private OpenAIClient openAiClient; + + @Mock + private ResponseService responseService; + + private ToolReference ref(String name, String description) { + return ToolReference.builder() + .toolName(name) + .summary(description) + .build(); + } + + private OpenAiToolIndex newToolIndex() { + return OpenAiToolIndex.builder() + .openAiClient(this.openAiClient) + .model(MODEL) + .build(); + } + + private ToolSearchResponse search(OpenAiToolIndex toolIndex, String query) { + return toolIndex.search(new ToolSearchRequest(SESSION, query, null, null)); + } + + private Response responseWithTools(FunctionTool... tools) { + ResponseToolSearchOutputItem toolSearchOutputItem = ResponseToolSearchOutputItem.builder() + .id("tool-search-output-1") + .callId("tool-search-call-1") + .execution(ResponseToolSearchOutputItem.Execution.SERVER) + .status(ResponseToolSearchOutputItem.Status.COMPLETED) + .tools(Stream.of(tools) + .map(Tool::ofFunction) + .toList()) + .build(); + + return Response.builder() + .id("response-1") + .createdAt(1_000.0) + .error(Optional.empty()) + .incompleteDetails(Optional.empty()) + .instructions(Optional.empty()) + .metadata(Optional.empty()) + .model(MODEL) + .output(List.of(ResponseOutputItem.ofToolSearchOutput(toolSearchOutputItem))) + .parallelToolCalls(false) + .temperature(1.0) + .toolChoice(ToolChoice.ofOptions(ToolChoiceOptions.AUTO)) + .tools(List.of()) + .topP(1.0) + .completedAt(1_001.25) + .build(); + } + + private FunctionTool functionTool(String name, String description) { + return FunctionTool.builder() + .name(name) + .description(description) + .parameters(Optional.empty()) + .strict(false) + .build(); + } + + @Test + void builderRequiresOpenAiClient() { + assertThatIllegalArgumentException().isThrownBy(() -> OpenAiToolIndex.builder().model(MODEL).build()) + .withMessage("openAiClient must not be null"); + } + + @Test + void builderRequiresModel() { + assertThatIllegalArgumentException() + .isThrownBy(() -> OpenAiToolIndex.builder().openAiClient(this.openAiClient).build()) + .withMessage("model must not be null"); + } + + @Test + void builderRejectsModelsBeforeToolSearchSupport() { + assertThatThrownBy(() -> OpenAiToolIndex.builder() + .openAiClient(this.openAiClient) + .model(ChatModel.GPT_4O.asString()) + .build()) + .isInstanceOf(OpenAIException.class) + .hasMessage("Only gpt-5.4 and later models support tool search."); + } + + @Test + void searchReturnsEmptyForSessionWithoutTools() { + try (OpenAiToolIndex toolIndex = newToolIndex()) { + ToolSearchResponse response = search(toolIndex, "weather"); + + assertThat(response.toolReferences()).isEmpty(); + verify(this.responseService, never()).create(any(ResponseCreateParams.class)); + } + } + + @Test + void searchSendsQueryAndIndexedToolsToOpenAi() { + when(this.openAiClient.responses()).thenReturn(this.responseService); + when(this.responseService.create(any(ResponseCreateParams.class))) + .thenReturn(responseWithTools(functionTool("weatherTool", "Returns current weather conditions"))); + try (OpenAiToolIndex toolIndex = newToolIndex()) { + toolIndex.indexTool(SESSION, ref("weatherTool", "Returns current weather conditions")); + + ToolSearchResponse response = search(toolIndex, "weather"); + + assertThat(response.toolReferences()).hasSize(1); + assertThat(response.toolReferences().get(0).toolName()).isEqualTo("weatherTool"); + assertThat(response.toolReferences().get(0).summary()).isEqualTo("Returns current weather conditions"); + + ArgumentCaptor captor = ArgumentCaptor.forClass(ResponseCreateParams.class); + verify(this.responseService).create(captor.capture()); + + ResponseCreateParams params = captor.getValue(); + assertThat(params.input()).hasValueSatisfying(input -> assertThat(input.asText()).isEqualTo("weather")); + assertThat(params.model()).hasValueSatisfying(model -> assertThat(model.asString()).isEqualTo(MODEL)); + assertThat(params.parallelToolCalls()).contains(false); + assertThat(params.tools()).hasValueSatisfying(tools -> { + assertThat(tools).hasSize(2); + assertThat(tools.get(0).isSearch()).isTrue(); + assertThat(tools.get(1).asFunction().name()).isEqualTo("weatherTool"); + assertThat(tools.get(1).asFunction().description()).contains("Returns current weather conditions"); + assertThat(tools.get(1).asFunction().parameters()).isEmpty(); + assertThat(tools.get(1).asFunction().strict()).contains(false); + }); + } + } + + @Test + void indexToolsBatchAddsAllToolsToOpenAiRequest() { + when(this.openAiClient.responses()).thenReturn(this.responseService); + when(this.responseService.create(any(ResponseCreateParams.class))).thenReturn(responseWithTools()); + try (OpenAiToolIndex toolIndex = newToolIndex()) { + toolIndex.indexTools(SESSION, + List.of(ref("tool1", "First tool"), ref("tool2", "Second tool"), ref("tool3", "Third tool"))); + + search(toolIndex, "tool"); + + ArgumentCaptor captor = ArgumentCaptor.forClass(ResponseCreateParams.class); + verify(this.responseService).create(captor.capture()); + + assertThat(captor.getValue().tools()).hasValueSatisfying(tools -> { + assertThat(tools).hasSize(4); + assertThat(tools.stream().filter(Tool::isFunction).map(tool -> tool.asFunction().name())) + .containsExactly("tool1", "tool2", "tool3"); + }); + } + } + + @Test + void clearIndexRemovesSessionTools() { + try (OpenAiToolIndex toolIndex = newToolIndex()) { + toolIndex.indexTools(SESSION, List.of(ref("tool1", "desc1"), ref("tool2", "desc2"))); + + toolIndex.clearIndex(SESSION); + + ToolSearchResponse response = search(toolIndex, "tool"); + assertThat(response.toolReferences()).isEmpty(); + verify(this.responseService, never()).create(any(ResponseCreateParams.class)); + } + } + + @Test + void sessionIsolationPreventsLeakage() { + when(this.openAiClient.responses()).thenReturn(this.responseService); + when(this.responseService.create(any(ResponseCreateParams.class))).thenReturn(responseWithTools()); + try (OpenAiToolIndex toolIndex = newToolIndex()) { + toolIndex.indexTool(SESSION, ref("weatherTool", "Weather data")); + toolIndex.indexTool("session-2", ref("calculatorTool", "Math operations")); + + toolIndex.search(new ToolSearchRequest("session-2", "calculator", null, null)); + + ArgumentCaptor captor = ArgumentCaptor.forClass(ResponseCreateParams.class); + verify(this.responseService).create(captor.capture()); + assertThat(captor.getValue().tools()).hasValueSatisfying(tools -> { + assertThat(tools).hasSize(2); + assertThat(tools.get(1).asFunction().name()).isEqualTo("calculatorTool"); + }); + } + } + + @Test + void searchMetadataContainsSearchTypeQueryAndElapsedTime() { + when(this.openAiClient.responses()).thenReturn(this.responseService); + when(this.responseService.create(any(ResponseCreateParams.class))) + .thenReturn(responseWithTools(functionTool("weatherTool", "Weather data"))); + try (OpenAiToolIndex toolIndex = newToolIndex()) { + toolIndex.indexTool(SESSION, ref("weatherTool", "Weather data")); + + ToolSearchResponse response = search(toolIndex, "weather"); + ToolSearchResponse.SearchMetadata searchMetadata = Objects.requireNonNull(response.searchMetadata()); + + assertThat(searchMetadata.searchType()).isEqualTo("OpenAiToolIndex"); + assertThat(searchMetadata.query()).isEqualTo("weather"); + assertThat(searchMetadata.searchTimeMs()).isEqualTo(1_250L); + } + } + + @Test + void closeClosesOpenAiClient() { + try (OpenAiToolIndex toolIndex = newToolIndex()) { + assertThat(toolIndex).isNotNull(); + } + + verify(this.openAiClient).close(); + } + +}