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(); + } + +}